1#![expect(missing_docs)]
7#![forbid(unsafe_code)]
8
9mod buffers;
10pub mod resolver;
11mod rx_bufs;
12mod saved_state;
13mod test;
14
15pub use netvsp_protocol::protocol;
19pub use netvsp_protocol::rndisprot;
20
21use crate::buffers::GuestBuffers;
22use crate::protocol::VMS_SWITCH_RSS_MAX_SEND_INDIRECTION_TABLE_ENTRIES;
23use crate::protocol::Version;
24use crate::rndisprot::NDIS_HASH_FUNCTION_MASK;
25use crate::rndisprot::NDIS_RSS_PARAM_FLAG_DISABLE_RSS;
26use async_trait::async_trait;
27pub use buffers::BufferPool;
28use buffers::sub_allocation_size_for_mtu;
29use futures::FutureExt;
30use futures::StreamExt;
31use futures::channel::mpsc;
32use futures::channel::mpsc::TrySendError;
33use futures_concurrency::future::Race;
34use guestmem::AccessError;
35use guestmem::GuestMemory;
36use guestmem::GuestMemoryError;
37use guestmem::MemoryRead;
38use guestmem::MemoryWrite;
39use guestmem::ranges::GuestMemoryView;
40use guestmem::ranges::PagedRange;
41use guestmem::ranges::PagedRanges;
42use guestmem::ranges::PagedRangesReader;
43use guid::Guid;
44use hvdef::hypercall::HvGuestOsId;
45use hvdef::hypercall::HvGuestOsMicrosoft;
46use hvdef::hypercall::HvGuestOsMicrosoftIds;
47use hvdef::hypercall::HvGuestOsOpenSourceType;
48use inspect::Inspect;
49use inspect::InspectMut;
50use inspect::SensitivityLevel;
51use inspect_counters::Counter;
52use inspect_counters::Histogram;
53use mesh::rpc::Rpc;
54use net_backend::Endpoint;
55use net_backend::EndpointAction;
56use net_backend::QueueConfig;
57use net_backend::RxId;
58use net_backend::TxError;
59use net_backend::TxId;
60use net_backend::TxSegment;
61use net_backend_resources::mac_address::MacAddress;
62use pal_async::timer::Instant;
63use pal_async::timer::PolledTimer;
64use ring::gparange::MultiPagedRangeIter;
65use rx_bufs::RxBuffers;
66use rx_bufs::SubAllocationInUse;
67use std::collections::VecDeque;
68use std::fmt::Debug;
69use std::future::pending;
70use std::mem::offset_of;
71use std::ops::Range;
72use std::sync::Arc;
73use std::sync::atomic::AtomicUsize;
74use std::sync::atomic::Ordering;
75use std::task::Poll;
76use std::time::Duration;
77use task_control::AsyncRun;
78use task_control::InspectTaskMut;
79use task_control::StopTask;
80use task_control::TaskControl;
81use thiserror::Error;
82use tracing::Instrument;
83use vmbus_async::queue;
84use vmbus_async::queue::ExternalDataError;
85use vmbus_async::queue::IncomingPacket;
86use vmbus_async::queue::Queue;
87use vmbus_channel::bus::OfferParams;
88use vmbus_channel::bus::OpenRequest;
89use vmbus_channel::channel::ChannelControl;
90use vmbus_channel::channel::ChannelOpenError;
91use vmbus_channel::channel::ChannelRestoreError;
92use vmbus_channel::channel::DeviceResources;
93use vmbus_channel::channel::RestoreControl;
94use vmbus_channel::channel::SaveRestoreVmbusDevice;
95use vmbus_channel::channel::VmbusDevice;
96use vmbus_channel::gpadl::GpadlId;
97use vmbus_channel::gpadl::GpadlMapView;
98use vmbus_channel::gpadl::GpadlView;
99use vmbus_channel::gpadl::UnknownGpadlId;
100use vmbus_channel::gpadl_ring::GpadlRingMem;
101use vmbus_channel::gpadl_ring::gpadl_channel;
102use vmbus_ring as ring;
103use vmbus_ring::OutgoingPacketType;
104use vmbus_ring::RingMem;
105use vmbus_ring::gparange::MultiPagedRangeBuf;
106use vmcore::save_restore::RestoreError;
107use vmcore::save_restore::SaveError;
108use vmcore::save_restore::SavedStateBlob;
109use vmcore::vm_task::VmTaskDriver;
110use vmcore::vm_task::VmTaskDriverSource;
111use zerocopy::FromBytes;
112use zerocopy::FromZeros;
113use zerocopy::Immutable;
114use zerocopy::IntoBytes;
115use zerocopy::KnownLayout;
116
117const MIN_CONTROL_RING_SIZE: usize = 144;
120
121const MIN_STATE_CHANGE_RING_SIZE: usize = 196;
124
125const VF_ASSOCIATION_TRANSACTION_ID: u64 = 0x8000000000000000;
127const SWITCH_DATA_PATH_TRANSACTION_ID: u64 = 0x8000000000000001;
129
130const NETVSP_MAX_SUBCHANNELS_PER_VNIC: u16 = 64;
131
132#[cfg(not(test))]
140const VF_DEVICE_DELAY: Duration = Duration::from_secs(1);
141#[cfg(test)]
142const VF_DEVICE_DELAY: Duration = Duration::from_millis(100);
143
144#[cfg(not(test))]
147const LINK_DELAY_DURATION: Duration = Duration::from_secs(5);
148#[cfg(test)]
149const LINK_DELAY_DURATION: Duration = Duration::from_millis(333);
150
151#[derive(Default, PartialEq)]
152struct CoordinatorMessageUpdateType {
153 guest_vf_state: bool,
156 filter_state: bool,
158}
159
160#[derive(PartialEq)]
161enum CoordinatorMessage {
162 Update(CoordinatorMessageUpdateType),
164 Restart { channel_idx: u16 },
168 StartTimer(Instant),
170}
171
172struct Worker<T: RingMem> {
173 channel_idx: u16,
174 target_vp: Option<u32>,
175 mem: GuestMemory,
176 channel: NetChannel<T>,
177 state: WorkerState,
178 coordinator_send: mpsc::Sender<CoordinatorMessage>,
179}
180
181struct NetQueue {
182 driver: VmTaskDriver,
183 queue_state: Option<QueueState>,
184}
185
186impl<T: RingMem + 'static + Sync> InspectTaskMut<Worker<T>> for NetQueue {
187 fn inspect_mut(&mut self, req: inspect::Request<'_>, worker: Option<&mut Worker<T>>) {
188 if worker.is_none() && self.queue_state.is_none() {
189 req.ignore();
190 return;
191 }
192
193 let mut resp = req.respond();
194 resp.field("driver", &self.driver);
195 if let Some(worker) = worker {
196 resp.field(
197 "protocol_state",
198 match &worker.state {
199 WorkerState::Init(None) => "version",
200 WorkerState::Init(Some(_)) => "init",
201 WorkerState::Ready(_) => "ready",
202 WorkerState::WaitingForCoordinator(_) => "waiting for coordinator",
203 },
204 )
205 .field("ring", &worker.channel.queue)
206 .field(
207 "can_use_ring_size_optimization",
208 worker.channel.can_use_ring_size_opt,
209 );
210
211 if let Some(state) = worker.state.ready() {
212 resp.field(
213 "outstanding_tx_packets",
214 state.state.pending_tx_packets.len() - state.state.free_tx_packets.len(),
215 )
216 .field("pending_rx_packets", state.state.pending_rx_packets.len())
217 .field(
218 "pending_tx_completions",
219 state.state.pending_tx_completions.len(),
220 )
221 .field("free_tx_packets", state.state.free_tx_packets.len())
222 .merge(&state.state.stats);
223 }
224
225 resp.field("packet_filter", worker.channel.packet_filter)
226 .field(
227 "packet_size",
228 match worker.channel.packet_size {
229 PacketSize::V1 => protocol::PACKET_SIZE_V1,
230 PacketSize::V61 => protocol::PACKET_SIZE_V61,
231 },
232 );
233 }
234
235 if let Some(queue_state) = &mut self.queue_state {
236 resp.field_mut("queue", &mut queue_state.queue)
237 .field("rx_buffers", queue_state.rx_buffer_range.id_range.len())
238 .field(
239 "rx_buffers_start",
240 queue_state.rx_buffer_range.id_range.start,
241 );
242 }
243 }
244}
245
246enum WorkerState {
247 Init(Option<InitState>),
248 Ready(ReadyState),
249 WaitingForCoordinator(Option<ReadyState>),
250}
251
252impl WorkerState {
253 fn ready(&self) -> Option<&ReadyState> {
254 match self {
255 Self::Ready(state) | Self::WaitingForCoordinator(Some(state)) => Some(state),
256 _ => None,
257 }
258 }
259
260 fn ready_mut(&mut self) -> Option<&mut ReadyState> {
261 match self {
262 Self::Ready(state) | Self::WaitingForCoordinator(Some(state)) => Some(state),
263 _ => None,
264 }
265 }
266}
267
268struct InitState {
269 version: Version,
270 ndis_config: Option<NdisConfig>,
271 ndis_version: Option<NdisVersion>,
272 recv_buffer: Option<ReceiveBuffer>,
273 send_buffer: Option<SendBuffer>,
274}
275
276#[derive(Copy, Clone, Debug, Inspect)]
277struct NdisVersion {
278 #[inspect(hex)]
279 major: u32,
280 #[inspect(hex)]
281 minor: u32,
282}
283
284#[derive(Copy, Clone, Debug, Inspect)]
285struct NdisConfig {
286 #[inspect(safe)]
287 mtu: u32,
288 #[inspect(safe)]
289 capabilities: protocol::NdisConfigCapabilities,
290}
291
292struct ReadyState {
293 buffers: Arc<ChannelBuffers>,
294 state: ActiveState,
295 data: ProcessingData,
296}
297
298impl ReadyState {
299 fn reset_tx_after_endpoint_stop(&mut self) {
308 let state = &mut self.state;
309
310 let pending_tx = state
313 .pending_tx_packets
314 .iter_mut()
315 .enumerate()
316 .filter_map(|(id, inflight)| {
317 if inflight.pending_packet_count > 0 {
318 inflight.pending_packet_count = 0;
319 Some(PendingTxCompletion {
320 transaction_id: inflight.transaction_id,
321 tx_id: Some(TxId(id as u32)),
322 status: protocol::Status::SUCCESS,
323 })
324 } else {
325 None
326 }
327 })
328 .collect::<Vec<_>>();
329 state.pending_tx_completions.extend(pending_tx);
330
331 self.data.tx_segments.clear();
334 self.data.tx_segments_sent = 0;
335 }
336}
337
338#[async_trait]
341pub trait VirtualFunction: Sync + Send {
342 async fn id(&self) -> Option<u32>;
346 async fn guest_ready_for_device(&mut self);
348 async fn wait_for_state_change(&mut self) -> Rpc<(), ()>;
351}
352
353struct Adapter {
354 driver: VmTaskDriver,
355 mac_address: MacAddress,
356 max_queues: u16,
357 indirection_table_size: u16,
358 offload_support: OffloadConfig,
359 ring_size_limit: AtomicUsize,
360 free_tx_packet_threshold: usize,
361 tx_fast_completions: bool,
362 adapter_index: u32,
363 get_guest_os_id: Option<Box<dyn Fn() -> HvGuestOsId + Send + Sync>>,
364 num_sub_channels_opened: AtomicUsize,
365 link_speed: u64,
366}
367
368struct QueueState {
369 queue: Box<dyn net_backend::Queue>,
370 pool: BufferPool,
371 rx_buffer_range: RxBufferRange,
372 target_vp_set: bool,
373}
374
375struct RxBufferRange {
376 id_range: Range<u32>,
377 remote_buffer_id_recv: Option<mpsc::UnboundedReceiver<u32>>,
378 remote_ranges: Arc<RxBufferRanges>,
379}
380
381impl RxBufferRange {
382 fn new(
383 ranges: Arc<RxBufferRanges>,
384 id_range: Range<u32>,
385 remote_buffer_id_recv: Option<mpsc::UnboundedReceiver<u32>>,
386 ) -> Self {
387 Self {
388 id_range,
389 remote_buffer_id_recv,
390 remote_ranges: ranges,
391 }
392 }
393
394 fn send_if_remote(&self, id: u32) -> bool {
395 if id < RX_RESERVED_CONTROL_BUFFERS || self.id_range.contains(&id) {
398 false
399 } else {
400 let i = (id - RX_RESERVED_CONTROL_BUFFERS) / self.remote_ranges.buffers_per_queue;
401 let i = (i as usize).min(self.remote_ranges.buffer_id_send.len() - 1);
405 let _ = self.remote_ranges.buffer_id_send[i].unbounded_send(id);
406 true
407 }
408 }
409}
410
411#[derive(Debug, Error)]
412enum RxBufferConfigError {
413 #[error("queue_count must be at least 1")]
414 ZeroQueueCount,
415 #[error("buffer_count must be >= RX_RESERVED_CONTROL_BUFFERS")]
416 InsufficientBuffers,
417 #[error("buffers_per_queue must not be 0")]
418 ZeroBuffersPerQueue,
419}
420
421struct RxBufferRanges {
422 buffers_per_queue: u32,
423 buffer_id_send: Vec<mpsc::UnboundedSender<u32>>,
424}
425
426impl RxBufferRanges {
427 fn validate_params(buffer_count: u32, queue_count: u32) -> Result<u32, WorkerError> {
429 if queue_count == 0 {
430 return Err(WorkerError::InvalidRxBufferConfig(
431 RxBufferConfigError::ZeroQueueCount,
432 ));
433 }
434 if buffer_count < RX_RESERVED_CONTROL_BUFFERS {
435 return Err(WorkerError::InvalidRxBufferConfig(
436 RxBufferConfigError::InsufficientBuffers,
437 ));
438 }
439 let buffers_per_queue = (buffer_count - RX_RESERVED_CONTROL_BUFFERS) / queue_count;
440 if buffers_per_queue == 0 {
441 return Err(WorkerError::InvalidRxBufferConfig(
442 RxBufferConfigError::ZeroBuffersPerQueue,
443 ));
444 }
445 Ok(buffers_per_queue)
446 }
447
448 fn new(
449 buffer_count: u32,
450 queue_count: u32,
451 ) -> Result<(Self, Vec<mpsc::UnboundedReceiver<u32>>), WorkerError> {
452 let buffers_per_queue = Self::validate_params(buffer_count, queue_count)?;
453 #[expect(clippy::disallowed_methods)] let (send, recv): (Vec<_>, Vec<_>) = (0..queue_count).map(|_| mpsc::unbounded()).unzip();
455 Ok((
456 Self {
457 buffers_per_queue,
458 buffer_id_send: send,
459 },
460 recv,
461 ))
462 }
463}
464
465struct RssState {
466 key: [u8; 40],
467 indirection_table: Vec<u16>,
468}
469
470struct NetChannel<T: RingMem> {
472 adapter: Arc<Adapter>,
473 queue: Queue<T>,
474 gpadl_map: GpadlMapView,
475 packet_size: PacketSize,
476 pending_send_size: usize,
477 restart: Option<CoordinatorMessage>,
478 can_use_ring_size_opt: bool,
479 packet_filter: u32,
480}
481
482#[derive(Debug, Copy, Clone, PartialEq)]
484enum PacketSize {
485 V1,
487 V61,
489}
490
491impl From<Version> for PacketSize {
492 fn from(v: Version) -> Self {
493 if v >= Version::V61 {
494 PacketSize::V61
495 } else {
496 PacketSize::V1
497 }
498 }
499}
500
501struct ProcessingData {
503 tx_segments: Vec<TxSegment>,
504 tx_segments_sent: usize,
505 tx_done: Box<[TxId]>,
506 rx_ready: Box<[RxId]>,
507 rx_done: Vec<RxId>,
508 transfer_pages: Vec<ring::TransferPageRange>,
509 external_data: MultiPagedRangeBuf,
510}
511
512impl ProcessingData {
513 fn new() -> Self {
514 Self {
515 tx_segments: Vec::new(),
516 tx_segments_sent: 0,
517 tx_done: vec![TxId(0); 8192].into(),
518 rx_ready: vec![RxId(0); RX_BATCH_SIZE].into(),
519 rx_done: Vec::with_capacity(RX_BATCH_SIZE),
520 transfer_pages: Vec::with_capacity(RX_BATCH_SIZE),
521 external_data: MultiPagedRangeBuf::new(),
522 }
523 }
524}
525
526#[derive(Debug, Inspect)]
529struct ChannelBuffers {
530 version: Version,
531 #[inspect(skip)]
532 mem: GuestMemory,
533 #[inspect(skip)]
534 recv_buffer: ReceiveBuffer,
535 #[inspect(skip)]
536 send_buffer: Option<SendBuffer>,
537 ndis_version: NdisVersion,
538 #[inspect(safe)]
539 ndis_config: NdisConfig,
540}
541
542#[derive(Copy, Clone, Debug)]
544struct ControlMessageId(u32);
545
546struct ActiveState {
548 primary: Option<PrimaryChannelState>,
549
550 pending_tx_packets: Vec<PendingTxPacket>,
551 free_tx_packets: Vec<TxId>,
552 pending_tx_completions: VecDeque<PendingTxCompletion>,
553 pending_rx_packets: VecDeque<RxId>,
554
555 rx_bufs: RxBuffers,
556
557 stats: QueueStats,
558}
559
560#[derive(Inspect, Default)]
561struct QueueStats {
562 tx_stalled: Counter,
563 rx_dropped_ring_full: Counter,
564 rx_dropped_filtered: Counter,
565 spurious_wakes: Counter,
566 rx_packets: Counter,
567 tx_packets: Counter,
568 tx_lso_packets: Counter,
569 tx_checksum_packets: Counter,
570 tx_vlan_packets: Counter,
571 rx_vlan_packets: Counter,
572 tx_invalid_lso_packets: Counter,
573 tx_packets_per_wake: Histogram<10>,
574 rx_packets_per_wake: Histogram<10>,
575}
576
577#[derive(Debug)]
578struct PendingTxCompletion {
579 transaction_id: u64,
580 tx_id: Option<TxId>,
581 status: protocol::Status,
582}
583
584#[derive(Clone, Copy)]
585enum PrimaryChannelGuestVfState {
586 Initializing,
588 Restoring(saved_state::GuestVfState),
590 Unavailable,
592 UnavailableFromAvailable,
594 UnavailableFromDataPathSwitchPending { to_guest: bool, id: Option<u64> },
596 UnavailableFromDataPathSwitched,
598 Available { vfid: u32 },
600 AvailableAdvertised,
602 Ready,
604 DataPathSwitchPending {
606 to_guest: bool,
607 id: Option<u64>,
608 result: Option<bool>,
609 },
610 DataPathSwitched,
612 DataPathSynthetic,
615}
616
617impl std::fmt::Display for PrimaryChannelGuestVfState {
618 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
619 match self {
620 PrimaryChannelGuestVfState::Initializing => write!(f, "initializing"),
621 PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::NoState) => {
622 write!(f, "restoring")
623 }
624 PrimaryChannelGuestVfState::Restoring(
625 saved_state::GuestVfState::AvailableAdvertised,
626 ) => write!(f, "restoring from guest notified of vfid"),
627 PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::Ready) => {
628 write!(f, "restoring from vf present")
629 }
630 PrimaryChannelGuestVfState::Restoring(
631 saved_state::GuestVfState::DataPathSwitchPending {
632 to_guest, result, ..
633 },
634 ) => {
635 write!(
636 f,
637 "restoring from client requested data path switch: to {} {}",
638 if *to_guest { "guest" } else { "synthetic" },
639 if let Some(result) = result {
640 if *result { "succeeded\"" } else { "failed\"" }
641 } else {
642 "in progress\""
643 }
644 )
645 }
646 PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::DataPathSwitched) => {
647 write!(f, "restoring from data path in guest")
648 }
649 PrimaryChannelGuestVfState::Unavailable => write!(f, "unavailable"),
650 PrimaryChannelGuestVfState::UnavailableFromAvailable => {
651 write!(f, "\"unavailable (previously available)\"")
652 }
653 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending { .. } => {
654 write!(f, "unavailable (previously switching data path)")
655 }
656 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched => {
657 write!(f, "\"unavailable (previously using guest VF)\"")
658 }
659 PrimaryChannelGuestVfState::Available { vfid } => write!(f, "available vfid: {}", vfid),
660 PrimaryChannelGuestVfState::AvailableAdvertised => {
661 write!(f, "\"available, guest notified\"")
662 }
663 PrimaryChannelGuestVfState::Ready => write!(f, "\"available and present in guest\""),
664 PrimaryChannelGuestVfState::DataPathSwitchPending {
665 to_guest, result, ..
666 } => {
667 write!(
668 f,
669 "\"switching to {} {}",
670 if *to_guest { "guest" } else { "synthetic" },
671 if let Some(result) = result {
672 if *result { "succeeded\"" } else { "failed\"" }
673 } else {
674 "in progress\""
675 }
676 )
677 }
678 PrimaryChannelGuestVfState::DataPathSwitched => {
679 write!(f, "\"available and data path switched\"")
680 }
681 PrimaryChannelGuestVfState::DataPathSynthetic => write!(
682 f,
683 "\"available but data path switched back to synthetic due to external state change\""
684 ),
685 }
686 }
687}
688
689impl Inspect for PrimaryChannelGuestVfState {
690 fn inspect(&self, req: inspect::Request<'_>) {
691 req.value(self.to_string());
692 }
693}
694
695#[derive(Debug)]
696enum PendingLinkAction {
697 Default,
698 Active(bool),
699 Delay(bool),
700}
701
702struct PrimaryChannelState {
703 guest_vf_state: PrimaryChannelGuestVfState,
704 is_data_path_switched: Option<bool>,
705 control_messages: VecDeque<ControlMessage>,
706 control_messages_len: usize,
707 free_control_buffers: Vec<ControlMessageId>,
708 rss_state: Option<RssState>,
709 requested_num_queues: u16,
710 rndis_state: RndisState,
711 offload_config: OffloadConfig,
712 pending_offload_change: bool,
713 tx_spread_sent: bool,
714 guest_link_up: bool,
715 pending_link_action: PendingLinkAction,
716 advertised_vf_serial_number: Option<u32>,
719}
720
721impl Inspect for PrimaryChannelState {
722 fn inspect(&self, req: inspect::Request<'_>) {
723 req.respond()
724 .sensitivity_field(
725 "guest_vf_state",
726 SensitivityLevel::Safe,
727 self.guest_vf_state,
728 )
729 .sensitivity_field(
730 "data_path_switched",
731 SensitivityLevel::Safe,
732 self.is_data_path_switched,
733 )
734 .sensitivity_field(
735 "pending_control_messages",
736 SensitivityLevel::Safe,
737 self.control_messages.len(),
738 )
739 .sensitivity_field(
740 "free_control_message_buffers",
741 SensitivityLevel::Safe,
742 self.free_control_buffers.len(),
743 )
744 .sensitivity_field(
745 "pending_offload_change",
746 SensitivityLevel::Safe,
747 self.pending_offload_change,
748 )
749 .sensitivity_field("rndis_state", SensitivityLevel::Safe, self.rndis_state)
750 .sensitivity_field(
751 "offload_config",
752 SensitivityLevel::Safe,
753 &self.offload_config,
754 )
755 .sensitivity_field(
756 "tx_spread_sent",
757 SensitivityLevel::Safe,
758 self.tx_spread_sent,
759 )
760 .sensitivity_field("guest_link_up", SensitivityLevel::Safe, self.guest_link_up)
761 .sensitivity_field(
762 "pending_link_action",
763 SensitivityLevel::Safe,
764 match &self.pending_link_action {
765 PendingLinkAction::Active(up) => format!("Active({:x?})", up),
766 PendingLinkAction::Delay(up) => format!("Delay({:x?})", up),
767 PendingLinkAction::Default => "None".to_string(),
768 },
769 );
770 }
771}
772
773#[derive(Debug, Inspect, Clone)]
774struct OffloadConfig {
775 #[inspect(safe)]
776 checksum_tx: ChecksumOffloadConfig,
777 #[inspect(safe)]
778 checksum_rx: ChecksumOffloadConfig,
779 #[inspect(safe)]
780 lso4: bool,
781 #[inspect(safe)]
782 lso6: bool,
783}
784
785impl OffloadConfig {
786 fn mask_to_supported(&mut self, supported: &OffloadConfig) {
787 self.checksum_tx.mask_to_supported(&supported.checksum_tx);
788 self.checksum_rx.mask_to_supported(&supported.checksum_rx);
789 self.lso4 &= supported.lso4;
790 self.lso6 &= supported.lso6;
791 }
792}
793
794#[derive(Debug, Inspect, Clone)]
795struct ChecksumOffloadConfig {
796 #[inspect(safe)]
797 ipv4_header: bool,
798 #[inspect(safe)]
799 tcp4: bool,
800 #[inspect(safe)]
801 udp4: bool,
802 #[inspect(safe)]
803 tcp6: bool,
804 #[inspect(safe)]
805 udp6: bool,
806}
807
808impl ChecksumOffloadConfig {
809 fn mask_to_supported(&mut self, supported: &ChecksumOffloadConfig) {
810 self.ipv4_header &= supported.ipv4_header;
811 self.tcp4 &= supported.tcp4;
812 self.udp4 &= supported.udp4;
813 self.tcp6 &= supported.tcp6;
814 self.udp6 &= supported.udp6;
815 }
816
817 fn flags(
818 &self,
819 ) -> (
820 rndisprot::Ipv4ChecksumOffload,
821 rndisprot::Ipv6ChecksumOffload,
822 ) {
823 let on = rndisprot::NDIS_OFFLOAD_SUPPORTED;
824 let mut v4 = rndisprot::Ipv4ChecksumOffload::new();
825 let mut v6 = rndisprot::Ipv6ChecksumOffload::new();
826 if self.ipv4_header {
827 v4.set_ip_options_supported(on);
828 v4.set_ip_checksum(on);
829 }
830 if self.tcp4 {
831 v4.set_ip_options_supported(on);
832 v4.set_tcp_options_supported(on);
833 v4.set_tcp_checksum(on);
834 }
835 if self.tcp6 {
836 v6.set_ip_extension_headers_supported(on);
837 v6.set_tcp_options_supported(on);
838 v6.set_tcp_checksum(on);
839 }
840 if self.udp4 {
841 v4.set_ip_options_supported(on);
842 v4.set_udp_checksum(on);
843 }
844 if self.udp6 {
845 v6.set_ip_extension_headers_supported(on);
846 v6.set_udp_checksum(on);
847 }
848 (v4, v6)
849 }
850}
851
852impl OffloadConfig {
853 fn ndis_offload(&self) -> rndisprot::NdisOffload {
854 let checksum = {
855 let (ipv4_tx_flags, ipv6_tx_flags) = self.checksum_tx.flags();
856 let (ipv4_rx_flags, ipv6_rx_flags) = self.checksum_rx.flags();
857 rndisprot::TcpIpChecksumOffload {
858 ipv4_tx_encapsulation: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
859 ipv4_tx_flags,
860 ipv4_rx_encapsulation: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
861 ipv4_rx_flags,
862 ipv6_tx_encapsulation: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
863 ipv6_tx_flags,
864 ipv6_rx_encapsulation: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
865 ipv6_rx_flags,
866 }
867 };
868
869 let lso_v2 = {
870 let mut lso = rndisprot::TcpLargeSendOffloadV2::new_zeroed();
871 if self.lso4 {
872 lso.ipv4_encapsulation = rndisprot::NDIS_ENCAPSULATION_IEEE_802_3;
873 lso.ipv4_max_offload_size = rndisprot::LSO_MAX_OFFLOAD_SIZE;
874 lso.ipv4_min_segment_count = rndisprot::LSO_MIN_SEGMENT_COUNT;
875 }
876 if self.lso6 {
877 lso.ipv6_encapsulation = rndisprot::NDIS_ENCAPSULATION_IEEE_802_3;
878 lso.ipv6_max_offload_size = rndisprot::LSO_MAX_OFFLOAD_SIZE;
879 lso.ipv6_min_segment_count = rndisprot::LSO_MIN_SEGMENT_COUNT;
880 lso.ipv6_flags = rndisprot::Ipv6LsoFlags::new()
881 .with_ip_extension_headers_supported(rndisprot::NDIS_OFFLOAD_SUPPORTED)
882 .with_tcp_options_supported(rndisprot::NDIS_OFFLOAD_SUPPORTED);
883 }
884 lso
885 };
886
887 rndisprot::NdisOffload {
888 header: rndisprot::NdisObjectHeader {
889 object_type: rndisprot::NdisObjectType::OFFLOAD,
890 revision: 3,
891 size: rndisprot::NDIS_SIZEOF_NDIS_OFFLOAD_REVISION_3 as u16,
892 },
893 checksum,
894 lso_v2,
895 ..FromZeros::new_zeroed()
896 }
897 }
898}
899
900#[derive(Debug, Inspect, PartialEq, Eq, Copy, Clone)]
901pub enum RndisState {
902 Initializing,
903 Operational,
904 Halted,
905}
906
907impl PrimaryChannelState {
908 fn new(offload_config: OffloadConfig) -> Self {
909 Self {
910 guest_vf_state: PrimaryChannelGuestVfState::Initializing,
911 is_data_path_switched: None,
912 control_messages: VecDeque::new(),
913 control_messages_len: 0,
914 free_control_buffers: (0..RX_RESERVED_CONTROL_BUFFERS)
915 .map(ControlMessageId)
916 .collect(),
917 rss_state: None,
918 requested_num_queues: 1,
919 rndis_state: RndisState::Initializing,
920 pending_offload_change: false,
921 offload_config,
922 tx_spread_sent: false,
923 guest_link_up: true,
924 pending_link_action: PendingLinkAction::Default,
925 advertised_vf_serial_number: None,
926 }
927 }
928
929 fn restore(
930 guest_vf_state: &saved_state::GuestVfState,
931 rndis_state: &saved_state::RndisState,
932 offload_config: &saved_state::OffloadConfig,
933 pending_offload_change: bool,
934 advertised_vf_serial_number: Option<u32>,
935 num_queues: u16,
936 indirection_table_size: u16,
937 rx_bufs: &RxBuffers,
938 control_messages: Vec<saved_state::IncomingControlMessage>,
939 rss_state: Option<saved_state::RssState>,
940 tx_spread_sent: bool,
941 guest_link_down: bool,
942 pending_link_action: Option<bool>,
943 ) -> Result<Self, NetRestoreError> {
944 let control_messages_len = control_messages.iter().map(|msg| msg.data.len()).sum();
946
947 let control_messages = control_messages
948 .into_iter()
949 .map(|msg| ControlMessage {
950 message_type: msg.message_type,
951 data: msg.data.into(),
952 })
953 .collect();
954
955 let free_control_buffers = (0..RX_RESERVED_CONTROL_BUFFERS)
957 .filter_map(|id| rx_bufs.is_free(id).then_some(ControlMessageId(id)))
958 .collect();
959
960 let rss_state = rss_state
961 .map(|mut rss| {
962 if rss.indirection_table.len() > indirection_table_size as usize {
963 return Err(NetRestoreError::ReducedIndirectionTableSize);
966 }
967 if rss.indirection_table.len() < indirection_table_size as usize {
968 tracing::warn!(
969 saved_indirection_table_size = rss.indirection_table.len(),
970 adapter_indirection_table_size = indirection_table_size,
971 "increasing indirection table size",
972 );
973 let table_clone = rss.indirection_table.clone();
976 let num_to_add = indirection_table_size as usize - rss.indirection_table.len();
977 rss.indirection_table
978 .extend(table_clone.iter().cycle().take(num_to_add));
979 }
980 Ok(RssState {
981 key: rss
982 .key
983 .try_into()
984 .map_err(|_| NetRestoreError::InvalidRssKeySize)?,
985 indirection_table: rss.indirection_table,
986 })
987 })
988 .transpose()?;
989
990 let rndis_state = match rndis_state {
991 saved_state::RndisState::Initializing => RndisState::Initializing,
992 saved_state::RndisState::Operational => RndisState::Operational,
993 saved_state::RndisState::Halted => RndisState::Halted,
994 };
995
996 let guest_vf_state = PrimaryChannelGuestVfState::Restoring(*guest_vf_state);
997 let offload_config = OffloadConfig {
998 checksum_tx: ChecksumOffloadConfig {
999 ipv4_header: offload_config.checksum_tx.ipv4_header,
1000 tcp4: offload_config.checksum_tx.tcp4,
1001 udp4: offload_config.checksum_tx.udp4,
1002 tcp6: offload_config.checksum_tx.tcp6,
1003 udp6: offload_config.checksum_tx.udp6,
1004 },
1005 checksum_rx: ChecksumOffloadConfig {
1006 ipv4_header: offload_config.checksum_rx.ipv4_header,
1007 tcp4: offload_config.checksum_rx.tcp4,
1008 udp4: offload_config.checksum_rx.udp4,
1009 tcp6: offload_config.checksum_rx.tcp6,
1010 udp6: offload_config.checksum_rx.udp6,
1011 },
1012 lso4: offload_config.lso4,
1013 lso6: offload_config.lso6,
1014 };
1015
1016 let pending_link_action = if let Some(pending) = pending_link_action {
1017 PendingLinkAction::Active(pending)
1018 } else {
1019 PendingLinkAction::Default
1020 };
1021
1022 Ok(Self {
1023 guest_vf_state,
1024 is_data_path_switched: None,
1025 control_messages,
1026 control_messages_len,
1027 free_control_buffers,
1028 rss_state,
1029 requested_num_queues: num_queues,
1030 rndis_state,
1031 pending_offload_change,
1032 offload_config,
1033 tx_spread_sent,
1034 guest_link_up: !guest_link_down,
1035 pending_link_action,
1036 advertised_vf_serial_number,
1037 })
1038 }
1039}
1040
1041struct ControlMessage {
1042 message_type: u32,
1043 data: Box<[u8]>,
1044}
1045
1046const TX_PACKET_QUOTA: usize = 1024;
1047
1048impl ActiveState {
1049 fn new(primary: Option<PrimaryChannelState>, recv_buffer_count: u32) -> Self {
1050 Self {
1051 primary,
1052 pending_tx_packets: vec![Default::default(); TX_PACKET_QUOTA],
1053 free_tx_packets: (0..TX_PACKET_QUOTA as u32).rev().map(TxId).collect(),
1054 pending_tx_completions: VecDeque::new(),
1055 pending_rx_packets: VecDeque::new(),
1056 rx_bufs: RxBuffers::new(recv_buffer_count),
1057 stats: Default::default(),
1058 }
1059 }
1060
1061 fn restore(
1062 channel: &saved_state::Channel,
1063 recv_buffer_count: u32,
1064 ) -> Result<Self, NetRestoreError> {
1065 let mut active = Self::new(None, recv_buffer_count);
1066 let saved_state::Channel {
1067 pending_tx_completions,
1068 in_use_rx,
1069 } = channel;
1070 for rx in in_use_rx {
1071 active
1072 .rx_bufs
1073 .allocate(rx.buffers.as_slice().iter().copied())?;
1074 }
1075 for &transaction_id in pending_tx_completions {
1076 let tx_id = active.free_tx_packets.pop();
1080 if let Some(id) = tx_id {
1081 active.pending_tx_packets[id.0 as usize].transaction_id = transaction_id;
1083 }
1084 active
1087 .pending_tx_completions
1088 .push_back(PendingTxCompletion {
1089 transaction_id,
1090 tx_id,
1091 status: protocol::Status::SUCCESS,
1092 });
1093 }
1094 Ok(active)
1095 }
1096}
1097
1098#[derive(Default, Clone)]
1101struct PendingTxPacket {
1102 pending_packet_count: usize,
1103 transaction_id: u64,
1104}
1105
1106const RX_BATCH_SIZE: usize = 375;
1111
1112const RX_RESERVED_CONTROL_BUFFERS: u32 = 16;
1114
1115pub struct Nic {
1117 instance_id: Guid,
1118 offer_order: Option<u64>,
1119 resources: DeviceResources,
1120 coordinator: TaskControl<CoordinatorState, Coordinator>,
1121 coordinator_send: Option<mpsc::Sender<CoordinatorMessage>>,
1122 adapter: Arc<Adapter>,
1123 driver_source: VmTaskDriverSource,
1124}
1125
1126pub struct NicBuilder {
1127 virtual_function: Option<Box<dyn VirtualFunction>>,
1128 offer_order: Option<u64>,
1129 limit_ring_buffer: bool,
1130 max_queues: u16,
1131 get_guest_os_id: Option<Box<dyn Fn() -> HvGuestOsId + Send + Sync>>,
1132}
1133
1134impl NicBuilder {
1135 pub fn limit_ring_buffer(mut self, limit: bool) -> Self {
1136 self.limit_ring_buffer = limit;
1137 self
1138 }
1139
1140 pub fn max_queues(mut self, max_queues: u16) -> Self {
1141 self.max_queues = max_queues;
1142 self
1143 }
1144
1145 pub fn virtual_function(mut self, virtual_function: Box<dyn VirtualFunction>) -> Self {
1146 self.virtual_function = Some(virtual_function);
1147 self
1148 }
1149
1150 pub fn offer_order(mut self, offer_order: u64) -> Self {
1156 self.offer_order = Some(offer_order);
1157 self
1158 }
1159
1160 pub fn get_guest_os_id(mut self, os_type: Box<dyn Fn() -> HvGuestOsId + Send + Sync>) -> Self {
1161 self.get_guest_os_id = Some(os_type);
1162 self
1163 }
1164
1165 pub fn build(
1167 self,
1168 driver_source: &VmTaskDriverSource,
1169 instance_id: Guid,
1170 endpoint: Box<dyn Endpoint>,
1171 mac_address: MacAddress,
1172 adapter_index: u32,
1173 ) -> Nic {
1174 let multiqueue = endpoint.multiqueue_support();
1175
1176 let max_queues = self.max_queues.clamp(
1177 1,
1178 multiqueue.max_queues.min(NETVSP_MAX_SUBCHANNELS_PER_VNIC),
1179 );
1180
1181 let ring_size_limit = if self.limit_ring_buffer { 1024 } else { 0 };
1186
1187 let free_tx_packet_threshold = if endpoint.tx_fast_completions() {
1192 TX_PACKET_QUOTA
1193 } else {
1194 TX_PACKET_QUOTA / 4
1197 };
1198
1199 let tx_offloads = endpoint.tx_offload_support();
1200
1201 let offload_support = OffloadConfig {
1204 checksum_rx: ChecksumOffloadConfig {
1205 ipv4_header: true,
1206 tcp4: true,
1207 udp4: true,
1208 tcp6: true,
1209 udp6: true,
1210 },
1211 checksum_tx: ChecksumOffloadConfig {
1212 ipv4_header: tx_offloads.ipv4_header,
1213 tcp4: tx_offloads.tcp,
1214 tcp6: tx_offloads.tcp,
1215 udp4: tx_offloads.udp,
1216 udp6: tx_offloads.udp,
1217 },
1218 lso4: tx_offloads.tso && tx_offloads.ipv4_header,
1222 lso6: tx_offloads.tso,
1223 };
1224
1225 let driver = driver_source.simple();
1226 let adapter = Arc::new(Adapter {
1227 driver,
1228 mac_address,
1229 max_queues,
1230 indirection_table_size: multiqueue.indirection_table_size,
1231 offload_support,
1232 free_tx_packet_threshold,
1233 ring_size_limit: ring_size_limit.into(),
1234 tx_fast_completions: endpoint.tx_fast_completions(),
1235 adapter_index,
1236 get_guest_os_id: self.get_guest_os_id,
1237 num_sub_channels_opened: AtomicUsize::new(0),
1238 link_speed: endpoint.link_speed(),
1239 });
1240
1241 let coordinator = TaskControl::new(CoordinatorState {
1242 endpoint,
1243 adapter: adapter.clone(),
1244 virtual_function: self.virtual_function,
1245 pending_vf_state: CoordinatorStatePendingVfState::Ready,
1246 });
1247
1248 Nic {
1249 instance_id,
1250 offer_order: self.offer_order,
1251 resources: Default::default(),
1252 coordinator,
1253 coordinator_send: None,
1254 adapter,
1255 driver_source: driver_source.clone(),
1256 }
1257 }
1258}
1259
1260fn can_use_ring_opt<T: RingMem>(queue: &mut Queue<T>, guest_os_id: Option<HvGuestOsId>) -> bool {
1261 let Some(guest_os_id) = guest_os_id else {
1262 return false;
1264 };
1265
1266 if !queue.split().0.supports_pending_send_size() {
1267 return false;
1269 }
1270
1271 let Some(open_source_os) = guest_os_id.open_source() else {
1272 return true;
1274 };
1275
1276 match HvGuestOsOpenSourceType(open_source_os.os_type()) {
1277 HvGuestOsOpenSourceType::FREEBSD => open_source_os.version() >= 1400097,
1280 HvGuestOsOpenSourceType::LINUX => {
1284 open_source_os.version() >= 199424
1287 }
1288 _ => true,
1289 }
1290}
1291
1292impl Nic {
1293 pub fn builder() -> NicBuilder {
1294 NicBuilder {
1295 virtual_function: None,
1296 offer_order: None,
1297 limit_ring_buffer: false,
1298 max_queues: !0,
1299 get_guest_os_id: None,
1300 }
1301 }
1302
1303 pub fn shutdown(self) -> (Box<dyn Endpoint>, MacAddress) {
1304 let (state, _) = self.coordinator.into_inner();
1305 (state.endpoint, self.adapter.mac_address)
1306 }
1307}
1308
1309impl InspectMut for Nic {
1310 fn inspect_mut(&mut self, req: inspect::Request<'_>) {
1311 self.coordinator.inspect_mut(req);
1312 }
1313}
1314
1315#[async_trait]
1316impl VmbusDevice for Nic {
1317 fn offer(&self) -> OfferParams {
1318 OfferParams {
1319 interface_name: "net".to_owned(),
1320 instance_id: self.instance_id,
1321 interface_id: Guid {
1322 data1: 0xf8615163,
1323 data2: 0xdf3e,
1324 data3: 0x46c5,
1325 data4: [0x91, 0x3f, 0xf2, 0xd2, 0xf9, 0x65, 0xed, 0xe],
1326 },
1327 subchannel_index: 0,
1328 offer_order: self.offer_order,
1329 mnf_interrupt_latency: Some(Duration::from_micros(100)),
1330 ..Default::default()
1331 }
1332 }
1333
1334 fn max_subchannels(&self) -> u16 {
1335 self.adapter.max_queues
1336 }
1337
1338 fn install(&mut self, resources: DeviceResources) {
1339 self.resources = resources;
1340 }
1341
1342 async fn open(
1343 &mut self,
1344 channel_idx: u16,
1345 open_request: &OpenRequest,
1346 ) -> Result<(), ChannelOpenError> {
1347 let state = if channel_idx == 0 {
1349 self.insert_coordinator(1, None);
1350 WorkerState::Init(None)
1351 } else {
1352 self.coordinator.stop().await;
1353 let buffers = self.coordinator.state().unwrap().buffers.clone().unwrap();
1355 WorkerState::Ready(ReadyState {
1356 state: ActiveState::new(None, buffers.recv_buffer.count),
1357 buffers,
1358 data: ProcessingData::new(),
1359 })
1360 };
1361
1362 let num_opened = self
1363 .adapter
1364 .num_sub_channels_opened
1365 .fetch_add(1, Ordering::SeqCst);
1366 let r = self.insert_worker(channel_idx, open_request, state, true);
1367 if channel_idx != 0
1368 && num_opened + 1 == self.coordinator.state_mut().unwrap().num_queues as usize
1369 {
1370 let coordinator = &mut self.coordinator.state_mut().unwrap();
1371 coordinator.workers[0].stop().await;
1372 }
1373
1374 if r.is_err() && channel_idx == 0 {
1375 self.coordinator.remove();
1376 } else {
1377 self.coordinator.start();
1379 }
1380 r?;
1381 Ok(())
1382 }
1383
1384 async fn close(&mut self, channel_idx: u16) {
1385 if !self.coordinator.has_state() {
1386 tracing::error!(
1387 channel_idx,
1388 instance_id = %self.instance_id,
1389 "Close called while vmbus channel is already closed"
1390 );
1391 return;
1392 }
1393
1394 let restart = self.coordinator.stop().await;
1396
1397 {
1399 let worker = &mut self.coordinator.state_mut().unwrap().workers[channel_idx as usize];
1400 worker.stop().await;
1401 if worker.has_state() {
1402 worker.remove();
1403 }
1404 }
1405
1406 self.adapter
1407 .num_sub_channels_opened
1408 .fetch_sub(1, Ordering::SeqCst);
1409 if channel_idx == 0 {
1411 for worker in &mut self.coordinator.state_mut().unwrap().workers {
1412 worker.task_mut().queue_state = None;
1413 }
1414
1415 self.coordinator
1417 .task_mut()
1418 .endpoint
1419 .stop()
1420 .instrument(tracing::info_span!(
1421 "stopping coordinator endpoint",
1422 instance_id = %self.instance_id,
1423 ))
1424 .await;
1425
1426 self.coordinator.remove();
1431 } else {
1432 if restart {
1434 self.coordinator.start();
1435 }
1436 }
1437 }
1438
1439 async fn retarget_vp(&mut self, channel_idx: u16, target_vp: u32) {
1440 if !self.coordinator.has_state() {
1441 return;
1442 }
1443
1444 let coordinator_running = self.coordinator.stop().await;
1446 let worker = &mut self.coordinator.state_mut().unwrap().workers[channel_idx as usize];
1447 worker.stop().await;
1448 let (net_queue, worker_state) = worker.get_mut();
1449
1450 net_queue.driver.retarget_vp(target_vp);
1452
1453 if let Some(worker_state) = worker_state {
1454 worker_state.target_vp = Some(target_vp);
1456 if let Some(queue_state) = &mut net_queue.queue_state {
1457 queue_state.target_vp_set = false;
1459 }
1460 }
1461
1462 if coordinator_running {
1464 self.coordinator.start();
1465 }
1466 }
1467
1468 fn start(&mut self) {
1469 if !self.coordinator.is_running() {
1470 self.coordinator.start();
1471 }
1472 }
1473
1474 async fn stop(&mut self) {
1475 self.coordinator.stop().await;
1476 if let Some(coordinator) = self.coordinator.state_mut() {
1477 coordinator.stop_workers().await;
1478 }
1479 }
1480
1481 fn supports_save_restore(&mut self) -> Option<&mut dyn SaveRestoreVmbusDevice> {
1482 Some(self)
1483 }
1484}
1485
1486#[async_trait]
1487impl SaveRestoreVmbusDevice for Nic {
1488 async fn save(&mut self) -> Result<SavedStateBlob, SaveError> {
1489 let state = self.saved_state();
1490 Ok(SavedStateBlob::new(state))
1491 }
1492
1493 async fn restore(
1494 &mut self,
1495 control: RestoreControl<'_>,
1496 state: SavedStateBlob,
1497 ) -> Result<(), RestoreError> {
1498 let state: saved_state::SavedState = state.parse()?;
1499 if let Err(err) = self.restore_state(control, state).await {
1500 tracing::error!(
1501 error = &err as &dyn std::error::Error,
1502 instance_id = %self.instance_id,
1503 "Failed restoring network vmbus state"
1504 );
1505 Err(err.into())
1506 } else {
1507 Ok(())
1508 }
1509 }
1510}
1511
1512impl Nic {
1513 fn insert_worker(
1517 &mut self,
1518 channel_idx: u16,
1519 open_request: &OpenRequest,
1520 state: WorkerState,
1521 start: bool,
1522 ) -> Result<(), OpenError> {
1523 let coordinator = self.coordinator.state_mut().unwrap();
1524
1525 let driver = coordinator.workers[channel_idx as usize]
1529 .task()
1530 .driver
1531 .clone();
1532 driver.retarget_vp(open_request.open_data.target_vp.unwrap_or_default());
1533
1534 let packet_size = match &state {
1535 WorkerState::Init(Some(init)) => init.version.into(),
1536 WorkerState::Ready(ready) | WorkerState::WaitingForCoordinator(Some(ready)) => {
1537 ready.buffers.version.into()
1538 }
1539 WorkerState::Init(None) | WorkerState::WaitingForCoordinator(None) => PacketSize::V1,
1540 };
1541
1542 let raw = gpadl_channel(&driver, &self.resources, open_request, channel_idx)
1543 .map_err(OpenError::Ring)?;
1544 let mut queue = Queue::new(raw).map_err(OpenError::Queue)?;
1545 let guest_os_id = self.adapter.get_guest_os_id.as_ref().map(|f| f());
1546 let can_use_ring_size_opt = can_use_ring_opt(&mut queue, guest_os_id);
1547 let worker = Worker {
1548 channel_idx,
1549 target_vp: open_request.open_data.target_vp,
1550 mem: self
1551 .resources
1552 .offer_resources
1553 .guest_memory(open_request)
1554 .clone(),
1555 channel: NetChannel {
1556 adapter: self.adapter.clone(),
1557 queue,
1558 gpadl_map: self.resources.gpadl_map.clone(),
1559 packet_size,
1560 pending_send_size: 0,
1561 restart: None,
1562 can_use_ring_size_opt,
1563 packet_filter: coordinator.active_packet_filter,
1564 },
1565 state,
1566 coordinator_send: self.coordinator_send.clone().unwrap(),
1567 };
1568 let instance_id = self.instance_id;
1569 let worker_task = &mut coordinator.workers[channel_idx as usize];
1570 worker_task.insert(
1571 driver,
1572 format!("netvsp-{}-{}", instance_id, channel_idx),
1573 worker,
1574 );
1575 if start {
1576 worker_task.start();
1577 }
1578 Ok(())
1579 }
1580}
1581
1582struct RestoreCoordinatorState {
1583 active_packet_filter: u32,
1584}
1585
1586impl Nic {
1587 fn insert_coordinator(&mut self, num_queues: u16, restoring: Option<RestoreCoordinatorState>) {
1589 let mut driver_builder = self.driver_source.builder();
1590 driver_builder.target_vp(0);
1593 driver_builder.run_on_target(!self.adapter.tx_fast_completions);
1598
1599 #[expect(clippy::disallowed_methods)] let (send, recv) = mpsc::channel(1);
1601 self.coordinator_send = Some(send);
1602 self.coordinator.insert(
1603 &self.adapter.driver,
1604 format!("netvsp-{}-coordinator", self.instance_id),
1605 Coordinator {
1606 recv,
1607 channel_control: self.resources.channel_control.clone(),
1608 restart: restoring.is_some(),
1609 workers: (0..self.adapter.max_queues)
1610 .map(|i| {
1611 TaskControl::new(NetQueue {
1612 queue_state: None,
1613 driver: driver_builder
1614 .build(format!("netvsp-{}-{}", self.instance_id, i)),
1615 })
1616 })
1617 .collect(),
1618 buffers: None,
1619 num_queues,
1620 active_packet_filter: restoring
1621 .map(|r| r.active_packet_filter)
1622 .unwrap_or(rndisprot::NDIS_PACKET_TYPE_NONE),
1623 sleep_deadline: None,
1624 },
1625 );
1626 }
1627}
1628
1629#[derive(Debug, Error)]
1630enum NetRestoreError {
1631 #[error("unsupported protocol version {0:#x}")]
1632 UnsupportedVersion(u32),
1633 #[error("send/receive buffer invalid gpadl ID")]
1634 UnknownGpadlId(#[from] UnknownGpadlId),
1635 #[error("failed to restore channels")]
1636 Channel(#[source] ChannelRestoreError),
1637 #[error(transparent)]
1638 ReceiveBuffer(#[from] BufferError),
1639 #[error(transparent)]
1640 SuballocationMisconfigured(#[from] SubAllocationInUse),
1641 #[error(transparent)]
1642 Open(#[from] OpenError),
1643 #[error("invalid rss key size")]
1644 InvalidRssKeySize,
1645 #[error("reduced indirection table size")]
1646 ReducedIndirectionTableSize,
1647}
1648
1649impl From<NetRestoreError> for RestoreError {
1650 fn from(err: NetRestoreError) -> Self {
1651 RestoreError::InvalidSavedState(anyhow::Error::new(err))
1652 }
1653}
1654
1655impl Nic {
1656 async fn restore_state(
1657 &mut self,
1658 mut control: RestoreControl<'_>,
1659 state: saved_state::SavedState,
1660 ) -> Result<(), NetRestoreError> {
1661 let mut saved_packet_filter = 0u32;
1662 if let Some(state) = state.open {
1663 let open = match &state.primary {
1664 saved_state::Primary::Version => vec![true],
1665 saved_state::Primary::Init(_) => vec![true],
1666 saved_state::Primary::Ready(ready) => {
1667 ready.channels.iter().map(|x| x.is_some()).collect()
1668 }
1669 };
1670
1671 let mut states: Vec<_> = open.iter().map(|_| None).collect();
1672
1673 let requests = control
1679 .restore(&open)
1680 .await
1681 .map_err(NetRestoreError::Channel)?;
1682
1683 match state.primary {
1684 saved_state::Primary::Version => {
1685 states[0] = Some(WorkerState::Init(None));
1686 }
1687 saved_state::Primary::Init(init) => {
1688 let version = check_version(init.version)
1689 .ok_or(NetRestoreError::UnsupportedVersion(init.version))?;
1690
1691 let recv_buffer = init
1692 .receive_buffer
1693 .map(|recv_buffer| {
1694 ReceiveBuffer::new(
1695 &self.resources.gpadl_map,
1696 recv_buffer.gpadl_id,
1697 recv_buffer.id,
1698 recv_buffer.sub_allocation_size,
1699 )
1700 })
1701 .transpose()?;
1702
1703 let send_buffer = init
1704 .send_buffer
1705 .map(|send_buffer| {
1706 SendBuffer::new(&self.resources.gpadl_map, send_buffer.gpadl_id)
1707 })
1708 .transpose()?;
1709
1710 let state = InitState {
1711 version,
1712 ndis_config: init.ndis_config.map(
1713 |saved_state::NdisConfig { mtu, capabilities }| NdisConfig {
1714 mtu,
1715 capabilities: capabilities.into(),
1716 },
1717 ),
1718 ndis_version: init.ndis_version.map(
1719 |saved_state::NdisVersion { major, minor }| NdisVersion {
1720 major,
1721 minor,
1722 },
1723 ),
1724 recv_buffer,
1725 send_buffer,
1726 };
1727 states[0] = Some(WorkerState::Init(Some(state)));
1728 }
1729 saved_state::Primary::Ready(ready) => {
1730 let saved_state::ReadyPrimary {
1731 version,
1732 receive_buffer,
1733 send_buffer,
1734 mut control_messages,
1735 mut rss_state,
1736 channels,
1737 ndis_version,
1738 ndis_config,
1739 rndis_state,
1740 guest_vf_state,
1741 offload_config,
1742 pending_offload_change,
1743 tx_spread_sent,
1744 guest_link_down,
1745 pending_link_action,
1746 packet_filter,
1747 advertised_vf_serial_number,
1748 } = ready;
1749
1750 saved_packet_filter = packet_filter.unwrap_or(rndisprot::NPROTO_PACKET_FILTER);
1752
1753 let version = check_version(version)
1754 .ok_or(NetRestoreError::UnsupportedVersion(version))?;
1755
1756 let request = requests[0].as_ref().unwrap();
1757 let buffers = Arc::new(ChannelBuffers {
1758 version,
1759 mem: self.resources.offer_resources.guest_memory(request).clone(),
1760 recv_buffer: ReceiveBuffer::new(
1761 &self.resources.gpadl_map,
1762 receive_buffer.gpadl_id,
1763 receive_buffer.id,
1764 receive_buffer.sub_allocation_size,
1765 )?,
1766 send_buffer: {
1767 if let Some(send_buffer) = send_buffer {
1768 Some(SendBuffer::new(
1769 &self.resources.gpadl_map,
1770 send_buffer.gpadl_id,
1771 )?)
1772 } else {
1773 None
1774 }
1775 },
1776 ndis_version: {
1777 let saved_state::NdisVersion { major, minor } = ndis_version;
1778 NdisVersion { major, minor }
1779 },
1780 ndis_config: {
1781 let saved_state::NdisConfig { mtu, capabilities } = ndis_config;
1782 NdisConfig {
1783 mtu,
1784 capabilities: capabilities.into(),
1785 }
1786 },
1787 });
1788
1789 for (channel_idx, channel) in channels.iter().enumerate() {
1790 let channel = if let Some(channel) = channel {
1791 channel
1792 } else {
1793 continue;
1794 };
1795
1796 let mut active = ActiveState::restore(channel, buffers.recv_buffer.count)?;
1797
1798 if channel_idx == 0 {
1800 let primary = PrimaryChannelState::restore(
1801 &guest_vf_state,
1802 &rndis_state,
1803 &offload_config,
1804 pending_offload_change,
1805 advertised_vf_serial_number,
1806 channels.len() as u16,
1807 self.adapter.indirection_table_size,
1808 &active.rx_bufs,
1809 std::mem::take(&mut control_messages),
1810 rss_state.take(),
1811 tx_spread_sent,
1812 guest_link_down,
1813 pending_link_action,
1814 )?;
1815 active.primary = Some(primary);
1816 }
1817
1818 states[channel_idx] = Some(WorkerState::Ready(ReadyState {
1819 buffers: buffers.clone(),
1820 state: active,
1821 data: ProcessingData::new(),
1822 }));
1823 }
1824 }
1825 }
1826
1827 self.insert_coordinator(
1830 states.len() as u16,
1831 Some(RestoreCoordinatorState {
1832 active_packet_filter: saved_packet_filter,
1833 }),
1834 );
1835
1836 for (channel_idx, (state, request)) in states.into_iter().zip(requests).enumerate() {
1837 if let Some(state) = state {
1838 self.insert_worker(channel_idx as u16, &request.unwrap(), state, false)?;
1839 }
1840 }
1841 } else {
1842 control
1843 .restore(&[false])
1844 .await
1845 .map_err(NetRestoreError::Channel)?;
1846 }
1847 Ok(())
1848 }
1849
1850 fn saved_state(&self) -> saved_state::SavedState {
1851 let open = if let Some(coordinator) = self.coordinator.state() {
1852 let primary = coordinator.workers[0].state().unwrap();
1853 let primary = match &primary.state {
1854 WorkerState::Init(None) => saved_state::Primary::Version,
1855 WorkerState::Init(Some(init)) => {
1856 saved_state::Primary::Init(saved_state::InitPrimary {
1857 version: init.version as u32,
1858 ndis_config: init.ndis_config.map(|NdisConfig { mtu, capabilities }| {
1859 saved_state::NdisConfig {
1860 mtu,
1861 capabilities: capabilities.into(),
1862 }
1863 }),
1864 ndis_version: init.ndis_version.map(|NdisVersion { major, minor }| {
1865 saved_state::NdisVersion { major, minor }
1866 }),
1867 receive_buffer: init.recv_buffer.as_ref().map(|x| x.saved_state()),
1868 send_buffer: init.send_buffer.as_ref().map(|x| saved_state::SendBuffer {
1869 gpadl_id: x.gpadl.id(),
1870 }),
1871 })
1872 }
1873 WorkerState::WaitingForCoordinator(Some(ready)) | WorkerState::Ready(ready) => {
1874 let primary = ready.state.primary.as_ref().unwrap();
1875
1876 let rndis_state = match primary.rndis_state {
1877 RndisState::Initializing => saved_state::RndisState::Initializing,
1878 RndisState::Operational => saved_state::RndisState::Operational,
1879 RndisState::Halted => saved_state::RndisState::Halted,
1880 };
1881
1882 let offload_config = saved_state::OffloadConfig {
1883 checksum_tx: saved_state::ChecksumOffloadConfig {
1884 ipv4_header: primary.offload_config.checksum_tx.ipv4_header,
1885 tcp4: primary.offload_config.checksum_tx.tcp4,
1886 udp4: primary.offload_config.checksum_tx.udp4,
1887 tcp6: primary.offload_config.checksum_tx.tcp6,
1888 udp6: primary.offload_config.checksum_tx.udp6,
1889 },
1890 checksum_rx: saved_state::ChecksumOffloadConfig {
1891 ipv4_header: primary.offload_config.checksum_rx.ipv4_header,
1892 tcp4: primary.offload_config.checksum_rx.tcp4,
1893 udp4: primary.offload_config.checksum_rx.udp4,
1894 tcp6: primary.offload_config.checksum_rx.tcp6,
1895 udp6: primary.offload_config.checksum_rx.udp6,
1896 },
1897 lso4: primary.offload_config.lso4,
1898 lso6: primary.offload_config.lso6,
1899 };
1900
1901 let control_messages = primary
1902 .control_messages
1903 .iter()
1904 .map(|message| saved_state::IncomingControlMessage {
1905 message_type: message.message_type,
1906 data: message.data.to_vec(),
1907 })
1908 .collect();
1909
1910 let rss_state = primary.rss_state.as_ref().map(|rss| saved_state::RssState {
1911 key: rss.key.into(),
1912 indirection_table: rss.indirection_table.clone(),
1913 });
1914
1915 let pending_link_action = match primary.pending_link_action {
1916 PendingLinkAction::Default => None,
1917 PendingLinkAction::Active(action) | PendingLinkAction::Delay(action) => {
1918 Some(action)
1919 }
1920 };
1921
1922 let channels = coordinator.workers[..coordinator.num_queues as usize]
1923 .iter()
1924 .map(|worker| {
1925 worker.state().map(|worker| {
1926 if let Some(ready) = worker.state.ready() {
1927 let pending_tx_completions = ready
1930 .state
1931 .pending_tx_completions
1932 .iter()
1933 .map(|pending| pending.transaction_id)
1934 .chain(ready.state.pending_tx_packets.iter().filter_map(
1935 |inflight| {
1936 (inflight.pending_packet_count > 0)
1937 .then_some(inflight.transaction_id)
1938 },
1939 ))
1940 .collect();
1941
1942 saved_state::Channel {
1943 pending_tx_completions,
1944 in_use_rx: {
1945 ready
1946 .state
1947 .rx_bufs
1948 .allocated()
1949 .map(|id| saved_state::Rx {
1950 buffers: id.collect(),
1951 })
1952 .collect()
1953 },
1954 }
1955 } else {
1956 saved_state::Channel {
1957 pending_tx_completions: Vec::new(),
1958 in_use_rx: Vec::new(),
1959 }
1960 }
1961 })
1962 })
1963 .collect();
1964
1965 let guest_vf_state = match primary.guest_vf_state {
1966 PrimaryChannelGuestVfState::Initializing
1967 | PrimaryChannelGuestVfState::Unavailable
1968 | PrimaryChannelGuestVfState::Available { .. } => {
1969 saved_state::GuestVfState::NoState
1970 }
1971 PrimaryChannelGuestVfState::UnavailableFromAvailable
1972 | PrimaryChannelGuestVfState::AvailableAdvertised => {
1973 saved_state::GuestVfState::AvailableAdvertised
1974 }
1975 PrimaryChannelGuestVfState::Ready => saved_state::GuestVfState::Ready,
1976 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending {
1977 to_guest,
1978 id,
1979 } => saved_state::GuestVfState::DataPathSwitchPending {
1980 to_guest,
1981 id,
1982 result: None,
1983 },
1984 PrimaryChannelGuestVfState::DataPathSwitchPending {
1985 to_guest,
1986 id,
1987 result,
1988 } => saved_state::GuestVfState::DataPathSwitchPending {
1989 to_guest,
1990 id,
1991 result,
1992 },
1993 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
1994 | PrimaryChannelGuestVfState::DataPathSwitched
1995 | PrimaryChannelGuestVfState::DataPathSynthetic => {
1996 saved_state::GuestVfState::DataPathSwitched
1997 }
1998 PrimaryChannelGuestVfState::Restoring(saved_state) => saved_state,
1999 };
2000
2001 let worker_0_packet_filter = coordinator.workers[0]
2002 .state()
2003 .unwrap()
2004 .channel
2005 .packet_filter;
2006 saved_state::Primary::Ready(saved_state::ReadyPrimary {
2007 version: ready.buffers.version as u32,
2008 receive_buffer: ready.buffers.recv_buffer.saved_state(),
2009 send_buffer: ready.buffers.send_buffer.as_ref().map(|sb| {
2010 saved_state::SendBuffer {
2011 gpadl_id: sb.gpadl.id(),
2012 }
2013 }),
2014 rndis_state,
2015 guest_vf_state,
2016 offload_config,
2017 pending_offload_change: primary.pending_offload_change,
2018 control_messages,
2019 rss_state,
2020 channels,
2021 ndis_config: {
2022 let NdisConfig { mtu, capabilities } = ready.buffers.ndis_config;
2023 saved_state::NdisConfig {
2024 mtu,
2025 capabilities: capabilities.into(),
2026 }
2027 },
2028 ndis_version: {
2029 let NdisVersion { major, minor } = ready.buffers.ndis_version;
2030 saved_state::NdisVersion { major, minor }
2031 },
2032 tx_spread_sent: primary.tx_spread_sent,
2033 guest_link_down: !primary.guest_link_up,
2034 pending_link_action,
2035 packet_filter: Some(worker_0_packet_filter),
2036 advertised_vf_serial_number: primary.advertised_vf_serial_number,
2037 })
2038 }
2039 WorkerState::WaitingForCoordinator(None) => {
2040 unreachable!("valid ready state")
2041 }
2042 };
2043
2044 let state = saved_state::OpenState { primary };
2045 Some(state)
2046 } else {
2047 None
2048 };
2049
2050 saved_state::SavedState { open }
2051 }
2052}
2053
2054#[derive(Debug, Error)]
2055enum MessageComponentError {
2056 #[error("header")]
2057 Header,
2058 #[error("per-packet information")]
2059 PerPacketInfo,
2060 #[error("message body")]
2061 Data,
2062 #[error("control message")]
2063 Control,
2064}
2065
2066#[derive(Debug, Error)]
2067enum WorkerError {
2068 #[error("packet error")]
2069 Packet(#[source] PacketError),
2070 #[error("unexpected packet order: {0}")]
2071 UnexpectedPacketOrder(#[source] PacketOrderError),
2072 #[error("unknown rndis message type: {0}")]
2073 UnknownRndisMessageType(u32),
2074 #[error("junk after rndis packet message: {0:#x}")]
2075 NonRndisPacketAfterPacket(u32),
2076 #[error("memory access error")]
2077 Access(#[from] AccessError),
2078 #[error("rndis message too small: {0}")]
2079 RndisMessageTooSmall(#[source] MessageComponentError),
2080 #[error("rndis headers missing or split across a page")]
2082 RndisBadHeaders,
2083 #[error("unsupported rndis behavior")]
2084 UnsupportedRndisBehavior,
2085 #[error("vmbus queue error")]
2086 Queue(#[from] queue::Error),
2087 #[error("too many control messages")]
2088 TooManyControlMessages,
2089 #[error("invalid rndis packet completion")]
2090 InvalidRndisPacketCompletion,
2091 #[error("missing transaction id")]
2092 MissingTransactionId,
2093 #[error("invalid gpadl")]
2094 InvalidGpadl(#[source] guestmem::InvalidGpn),
2095 #[error("guest buffers error")]
2096 GuestBuffers(#[source] buffers::GuestBuffersError),
2097 #[error("gpa direct error")]
2098 GpaDirectError(#[source] GuestMemoryError),
2099 #[error("endpoint")]
2100 Endpoint(#[source] anyhow::Error),
2101 #[error("message not supported on sub channel: {0}")]
2102 NotSupportedOnSubChannel(u32),
2103 #[error("the ring buffer ran out of space, which should not be possible")]
2104 OutOfSpace,
2105 #[error("send/receive buffer error")]
2106 Buffer(#[from] BufferError),
2107 #[error("invalid rndis state")]
2108 InvalidRndisState,
2109 #[error("rndis message type not implemented")]
2110 RndisMessageTypeNotImplemented,
2111 #[error("invalid TCP header offset {0}")]
2112 InvalidTcpHeaderOffset(u16),
2113 #[error("cancelled")]
2114 Cancelled(task_control::Cancelled),
2115 #[error("tearing down because send/receive buffer is revoked")]
2116 BufferRevoked,
2117 #[error("endpoint requires queue restart: {0}")]
2118 EndpointRequiresQueueRestart(#[source] anyhow::Error),
2119 #[error("Failed to send message to coordinator")]
2120 CoordinatorMessageSendFailed(#[source] TrySendError<CoordinatorMessage>),
2121 #[error("invalid rx buffer configuration: {0}")]
2122 InvalidRxBufferConfig(#[source] RxBufferConfigError),
2123}
2124
2125impl From<task_control::Cancelled> for WorkerError {
2126 fn from(value: task_control::Cancelled) -> Self {
2127 Self::Cancelled(value)
2128 }
2129}
2130
2131#[derive(Debug, Error)]
2132enum OpenError {
2133 #[error("error establishing ring buffer")]
2134 Ring(#[source] vmbus_channel::gpadl_ring::Error),
2135 #[error("error establishing vmbus queue")]
2136 Queue(#[source] queue::Error),
2137}
2138
2139#[derive(Debug, Error)]
2140enum PacketError {
2141 #[error("UnknownType {0}")]
2142 UnknownType(u32),
2143 #[error("Access")]
2144 Access(#[source] AccessError),
2145 #[error("ExternalData")]
2146 ExternalData(#[source] ExternalDataError),
2147 #[error("InvalidSendBufferIndex")]
2148 InvalidSendBufferIndex,
2149}
2150
2151#[derive(Debug, Error)]
2152enum PacketOrderError {
2153 #[error("Invalid PacketData")]
2154 InvalidPacketData,
2155 #[error("Unexpected RndisPacket")]
2156 UnexpectedRndisPacket,
2157 #[error("SendNdisVersion already exists")]
2158 SendNdisVersionExists,
2159 #[error("SendNdisConfig already exists")]
2160 SendNdisConfigExists,
2161 #[error("SendReceiveBuffer already exists")]
2162 SendReceiveBufferExists,
2163 #[error("SendReceiveBuffer missing MTU")]
2164 SendReceiveBufferMissingMTU,
2165 #[error("SendSendBuffer already exists")]
2166 SendSendBufferExists,
2167 #[error("SwitchDataPathCompletion during PrimaryChannelState")]
2168 SwitchDataPathCompletionPrimaryChannelState,
2169}
2170
2171#[derive(Debug)]
2172enum PacketData {
2173 Init(protocol::MessageInit),
2174 SendNdisVersion(protocol::Message1SendNdisVersion),
2175 SendReceiveBuffer(protocol::Message1SendReceiveBuffer),
2176 SendSendBuffer(protocol::Message1SendSendBuffer),
2177 RevokeReceiveBuffer(protocol::Message1RevokeReceiveBuffer),
2178 RevokeSendBuffer(protocol::Message1RevokeSendBuffer),
2179 RndisPacket(protocol::Message1SendRndisPacket),
2180 RndisPacketComplete(protocol::Message1SendRndisPacketComplete),
2181 SendNdisConfig(protocol::Message2SendNdisConfig),
2182 SwitchDataPath(protocol::Message4SwitchDataPath),
2183 OidQueryEx(protocol::Message5OidQueryEx),
2184 SubChannelRequest(protocol::Message5SubchannelRequest),
2185 SendVfAssociationCompletion,
2186 SwitchDataPathCompletion,
2187}
2188
2189#[derive(Debug)]
2190struct Packet<'a> {
2191 data: PacketData,
2192 transaction_id: Option<u64>,
2193 external_data: &'a MultiPagedRangeBuf,
2194}
2195
2196type PacketReader<'a> = PagedRangesReader<'a, MultiPagedRangeIter<'a>>;
2197
2198impl Packet<'_> {
2199 fn rndis_reader<'a>(&'a self, mem: &'a GuestMemory) -> PacketReader<'a> {
2200 PagedRanges::new(self.external_data.iter()).reader(mem)
2201 }
2202}
2203
2204fn read_packet_data<T: IntoBytes + FromBytes + Immutable + KnownLayout>(
2205 reader: &mut impl MemoryRead,
2206) -> Result<T, PacketError> {
2207 reader
2208 .read_plain()
2209 .map_err(PacketError::Access)
2210 .inspect_err(|_| tracelimit::info_ratelimited!("read_packet_data"))
2211}
2212
2213fn parse_packet<'a, T: RingMem>(
2214 packet_ref: &queue::PacketRef<'_, T>,
2215 send_buffer: Option<&SendBuffer>,
2216 version: Option<Version>,
2217 external_data: &'a mut MultiPagedRangeBuf,
2218) -> Result<Packet<'a>, PacketError> {
2219 external_data.clear();
2220 let packet = match packet_ref.as_ref() {
2221 IncomingPacket::Data(data) => data,
2222 IncomingPacket::Completion(completion) => {
2223 let data = if completion.transaction_id() == VF_ASSOCIATION_TRANSACTION_ID {
2224 PacketData::SendVfAssociationCompletion
2225 } else if completion.transaction_id() == SWITCH_DATA_PATH_TRANSACTION_ID {
2226 PacketData::SwitchDataPathCompletion
2227 } else {
2228 let mut reader = completion.reader();
2229 let header: protocol::MessageHeader = reader
2230 .read_plain()
2231 .map_err(PacketError::Access)
2232 .inspect_err(|_| {
2233 tracelimit::info_ratelimited!(
2234 tx_id = completion.transaction_id(),
2235 "parsing completion header"
2236 )
2237 })?;
2238 match header.message_type {
2239 protocol::MESSAGE1_TYPE_SEND_RNDIS_PACKET_COMPLETE => {
2240 PacketData::RndisPacketComplete(read_packet_data(&mut reader)?)
2241 }
2242 typ => return Err(PacketError::UnknownType(typ)),
2243 }
2244 };
2245 return Ok(Packet {
2246 data,
2247 transaction_id: Some(completion.transaction_id()),
2248 external_data,
2249 });
2250 }
2251 };
2252
2253 let mut reader = packet.reader();
2254 let header: protocol::MessageHeader = reader
2255 .read_plain()
2256 .map_err(PacketError::Access)
2257 .inspect_err(|_| tracelimit::info_ratelimited!("parsing data packet header"))?;
2258 let data = match header.message_type {
2259 protocol::MESSAGE_TYPE_INIT => PacketData::Init(read_packet_data(&mut reader)?),
2260 protocol::MESSAGE1_TYPE_SEND_NDIS_VERSION if version >= Some(Version::V1) => {
2261 PacketData::SendNdisVersion(read_packet_data(&mut reader)?)
2262 }
2263 protocol::MESSAGE1_TYPE_SEND_RECEIVE_BUFFER if version >= Some(Version::V1) => {
2264 PacketData::SendReceiveBuffer(read_packet_data(&mut reader)?)
2265 }
2266 protocol::MESSAGE1_TYPE_REVOKE_RECEIVE_BUFFER if version >= Some(Version::V1) => {
2267 PacketData::RevokeReceiveBuffer(read_packet_data(&mut reader)?)
2268 }
2269 protocol::MESSAGE1_TYPE_SEND_SEND_BUFFER if version >= Some(Version::V1) => {
2270 PacketData::SendSendBuffer(read_packet_data(&mut reader)?)
2271 }
2272 protocol::MESSAGE1_TYPE_REVOKE_SEND_BUFFER if version >= Some(Version::V1) => {
2273 PacketData::RevokeSendBuffer(read_packet_data(&mut reader)?)
2274 }
2275 protocol::MESSAGE1_TYPE_SEND_RNDIS_PACKET if version >= Some(Version::V1) => {
2276 let message: protocol::Message1SendRndisPacket = read_packet_data(&mut reader)?;
2277 if message.send_buffer_section_index != 0xffffffff {
2278 let send_buffer_suballocation = send_buffer
2279 .ok_or(PacketError::InvalidSendBufferIndex)?
2280 .gpadl
2281 .first()
2282 .unwrap()
2283 .try_subrange(
2284 message.send_buffer_section_index as usize * 6144,
2285 message.send_buffer_section_size as usize,
2286 )
2287 .ok_or(PacketError::InvalidSendBufferIndex)?;
2288
2289 external_data.push_range(send_buffer_suballocation);
2290 }
2291 PacketData::RndisPacket(message)
2292 }
2293 protocol::MESSAGE2_TYPE_SEND_NDIS_CONFIG if version >= Some(Version::V2) => {
2294 PacketData::SendNdisConfig(read_packet_data(&mut reader)?)
2295 }
2296 protocol::MESSAGE4_TYPE_SWITCH_DATA_PATH if version >= Some(Version::V4) => {
2297 PacketData::SwitchDataPath(read_packet_data(&mut reader)?)
2298 }
2299 protocol::MESSAGE5_TYPE_OID_QUERY_EX if version >= Some(Version::V5) => {
2300 PacketData::OidQueryEx(read_packet_data(&mut reader)?)
2301 }
2302 protocol::MESSAGE5_TYPE_SUB_CHANNEL if version >= Some(Version::V5) => {
2303 PacketData::SubChannelRequest(read_packet_data(&mut reader)?)
2304 }
2305 typ => return Err(PacketError::UnknownType(typ)),
2306 };
2307 packet
2308 .read_external_ranges(external_data)
2309 .map_err(PacketError::ExternalData)?;
2310 Ok(Packet {
2311 data,
2312 transaction_id: packet.transaction_id(),
2313 external_data,
2314 })
2315}
2316
2317#[derive(Debug, Copy, Clone)]
2318struct NvspMessage {
2319 buf: [u64; protocol::PACKET_SIZE_V61 / 8],
2320 size: PacketSize,
2321}
2322
2323impl NvspMessage {
2324 fn new<P: IntoBytes + Immutable + KnownLayout>(
2325 size: PacketSize,
2326 message_type: u32,
2327 data: P,
2328 ) -> Self {
2329 const {
2336 assert!(
2337 size_of::<P>() <= protocol::PACKET_SIZE_V1 - size_of::<protocol::MessageHeader>(),
2338 "packet might not fit in message"
2339 )
2340 };
2341 let mut message = NvspMessage {
2342 buf: [0; protocol::PACKET_SIZE_V61 / 8],
2343 size,
2344 };
2345 let header = protocol::MessageHeader { message_type };
2346 header.write_to_prefix(message.buf.as_mut_bytes()).unwrap();
2347 data.write_to_prefix(
2348 &mut message.buf.as_mut_bytes()[size_of::<protocol::MessageHeader>()..],
2349 )
2350 .unwrap();
2351 message
2352 }
2353
2354 fn aligned_payload(&self) -> &[u64] {
2355 let len = match self.size {
2358 PacketSize::V1 => const { protocol::PACKET_SIZE_V1.div_ceil(8) },
2359 PacketSize::V61 => const { protocol::PACKET_SIZE_V61.div_ceil(8) },
2360 };
2361 &self.buf[..len]
2362 }
2363}
2364
2365impl<T: RingMem> NetChannel<T> {
2366 fn message<P: IntoBytes + Immutable + KnownLayout>(
2367 &self,
2368 message_type: u32,
2369 data: P,
2370 ) -> NvspMessage {
2371 NvspMessage::new(self.packet_size, message_type, data)
2372 }
2373
2374 fn send_completion(
2375 &mut self,
2376 transaction_id: Option<u64>,
2377 message: Option<&NvspMessage>,
2378 ) -> Result<(), WorkerError> {
2379 match transaction_id {
2380 None => Ok(()),
2381 Some(transaction_id) => Ok(self
2382 .queue
2383 .split()
2384 .1
2385 .batched()
2386 .try_write_aligned(
2387 transaction_id,
2388 OutgoingPacketType::Completion,
2389 message.map_or(&[], |m| m.aligned_payload()),
2390 )
2391 .map_err(|err| match err {
2392 queue::TryWriteError::Full(_) => WorkerError::OutOfSpace,
2393 queue::TryWriteError::Queue(err) => WorkerError::Queue(err),
2394 })?),
2395 }
2396 }
2397}
2398
2399static SUPPORTED_VERSIONS: &[Version] = &[
2400 Version::V1,
2401 Version::V2,
2402 Version::V4,
2403 Version::V5,
2404 Version::V6,
2405 Version::V61,
2406];
2407
2408fn check_version(requested_version: u32) -> Option<Version> {
2409 SUPPORTED_VERSIONS
2410 .iter()
2411 .find(|version| **version as u32 == requested_version)
2412 .copied()
2413}
2414
2415#[derive(Debug)]
2416struct ReceiveBuffer {
2417 gpadl: GpadlView,
2418 id: u16,
2419 count: u32,
2420 sub_allocation_size: u32,
2421}
2422
2423#[derive(Debug, Error)]
2424enum BufferError {
2425 #[error("unsupported suballocation size {0}")]
2426 UnsupportedSuballocationSize(u32),
2427 #[error("unaligned gpadl")]
2428 UnalignedGpadl,
2429 #[error("unknown gpadl ID")]
2430 UnknownGpadlId(#[from] UnknownGpadlId),
2431}
2432
2433impl ReceiveBuffer {
2434 fn new(
2435 gpadl_map: &GpadlMapView,
2436 gpadl_id: GpadlId,
2437 id: u16,
2438 sub_allocation_size: u32,
2439 ) -> Result<Self, BufferError> {
2440 if sub_allocation_size < sub_allocation_size_for_mtu(DEFAULT_MTU) {
2441 return Err(BufferError::UnsupportedSuballocationSize(
2442 sub_allocation_size,
2443 ));
2444 }
2445 let gpadl = gpadl_map.map(gpadl_id)?;
2446 let range = gpadl
2447 .contiguous_aligned()
2448 .ok_or(BufferError::UnalignedGpadl)?;
2449 let num_sub_allocations = range.len() as u32 / sub_allocation_size;
2450 if num_sub_allocations == 0 {
2451 return Err(BufferError::UnsupportedSuballocationSize(
2452 sub_allocation_size,
2453 ));
2454 }
2455 let recv_buffer = Self {
2456 gpadl,
2457 id,
2458 count: num_sub_allocations,
2459 sub_allocation_size,
2460 };
2461 Ok(recv_buffer)
2462 }
2463
2464 fn range(&self, index: u32) -> PagedRange<'_> {
2465 self.gpadl.first().unwrap().subrange(
2466 (index * self.sub_allocation_size) as usize,
2467 self.sub_allocation_size as usize,
2468 )
2469 }
2470
2471 fn transfer_page_range(&self, index: u32, len: usize) -> ring::TransferPageRange {
2472 assert!(len <= self.sub_allocation_size as usize);
2473 ring::TransferPageRange {
2474 byte_offset: index * self.sub_allocation_size,
2475 byte_count: len as u32,
2476 }
2477 }
2478
2479 fn saved_state(&self) -> saved_state::ReceiveBuffer {
2480 saved_state::ReceiveBuffer {
2481 gpadl_id: self.gpadl.id(),
2482 id: self.id,
2483 sub_allocation_size: self.sub_allocation_size,
2484 }
2485 }
2486}
2487
2488#[derive(Debug)]
2489struct SendBuffer {
2490 gpadl: GpadlView,
2491}
2492
2493impl SendBuffer {
2494 fn new(gpadl_map: &GpadlMapView, gpadl_id: GpadlId) -> Result<Self, BufferError> {
2495 let gpadl = gpadl_map.map(gpadl_id)?;
2496 gpadl
2497 .contiguous_aligned()
2498 .ok_or(BufferError::UnalignedGpadl)?;
2499 Ok(Self { gpadl })
2500 }
2501}
2502
2503impl<T: RingMem> NetChannel<T> {
2504 fn handle_rndis_message(
2506 &mut self,
2507 state: &mut ActiveState,
2508 message_type: u32,
2509 mut reader: PacketReader<'_>,
2510 ) -> Result<(), WorkerError> {
2511 assert_ne!(
2512 message_type,
2513 rndisprot::MESSAGE_TYPE_PACKET_MSG,
2514 "handled elsewhere"
2515 );
2516 let control = state
2517 .primary
2518 .as_mut()
2519 .ok_or(WorkerError::NotSupportedOnSubChannel(message_type))?;
2520
2521 if message_type == rndisprot::MESSAGE_TYPE_HALT_MSG {
2522 return Ok(());
2524 }
2525
2526 const CONTROL_MESSAGE_MAX_QUEUED_BYTES: usize = 100 * 1024;
2531 if reader.len() == 0 {
2532 return Err(WorkerError::RndisMessageTooSmall(
2533 MessageComponentError::Control,
2534 ));
2535 }
2536 if CONTROL_MESSAGE_MAX_QUEUED_BYTES - control.control_messages_len < reader.len() {
2539 return Err(WorkerError::TooManyControlMessages);
2540 }
2541
2542 control.control_messages_len += reader.len();
2543 control.control_messages.push_back(ControlMessage {
2544 message_type,
2545 data: reader.read_all()?.into(),
2546 });
2547
2548 Ok(())
2551 }
2552
2553 fn handle_rndis_packet_messages(
2559 &mut self,
2560 buffers: &ChannelBuffers,
2561 state: &mut ActiveState,
2562 id: TxId,
2563 mut message_len: usize,
2564 mut reader: PacketReader<'_>,
2565 segments: &mut Vec<TxSegment>,
2566 ) -> Result<usize, WorkerError> {
2567 let mut num_packets = 0;
2571 loop {
2572 let next_message_offset = message_len
2573 .checked_sub(size_of::<rndisprot::MessageHeader>())
2574 .ok_or(WorkerError::RndisMessageTooSmall(
2575 MessageComponentError::Header,
2576 ))?;
2577
2578 self.handle_rndis_packet_message(
2579 id,
2580 reader.clone(),
2581 &buffers.mem,
2582 segments,
2583 &mut state.stats,
2584 )?;
2585 num_packets += 1;
2586
2587 reader.skip(next_message_offset)?;
2588 if reader.len() == 0 {
2589 break;
2590 }
2591 let header: rndisprot::MessageHeader = reader.read_plain()?;
2592 if header.message_type != rndisprot::MESSAGE_TYPE_PACKET_MSG {
2593 return Err(WorkerError::NonRndisPacketAfterPacket(header.message_type));
2594 }
2595 message_len = header.message_length as usize;
2596 }
2597 Ok(num_packets)
2598 }
2599
2600 fn handle_rndis_packet_message(
2602 &mut self,
2603 id: TxId,
2604 reader: PacketReader<'_>,
2605 mem: &GuestMemory,
2606 segments: &mut Vec<TxSegment>,
2607 stats: &mut QueueStats,
2608 ) -> Result<(), WorkerError> {
2609 let headers = reader
2611 .clone()
2612 .into_inner()
2613 .paged_ranges()
2614 .next()
2615 .ok_or(WorkerError::RndisBadHeaders)?;
2616 let mut data = reader.into_inner();
2617 let request: rndisprot::Packet = headers.reader(mem).read_plain()?;
2618 if request.num_oob_data_elements != 0
2619 || request.oob_data_length != 0
2620 || request.oob_data_offset != 0
2621 || request.vc_handle != 0
2622 {
2623 return Err(WorkerError::UnsupportedRndisBehavior);
2624 }
2625
2626 if data.len() < request.data_offset as usize
2627 || (data.len() - request.data_offset as usize) < request.data_length as usize
2628 || request.data_length == 0
2629 {
2630 return Err(WorkerError::RndisMessageTooSmall(
2631 MessageComponentError::Data,
2632 ));
2633 }
2634
2635 data.skip(request.data_offset as usize);
2636 data.truncate(request.data_length as usize);
2637
2638 let mut metadata = net_backend::TxMetadata {
2639 id,
2640 len: request.data_length,
2641 ..Default::default()
2642 };
2643
2644 if request.per_packet_info_length != 0 {
2645 let mut ppi = headers
2646 .try_subrange(
2647 request.per_packet_info_offset as usize,
2648 request.per_packet_info_length as usize,
2649 )
2650 .ok_or(WorkerError::RndisMessageTooSmall(
2651 MessageComponentError::PerPacketInfo,
2652 ))?;
2653 while !ppi.is_empty() {
2654 let h: rndisprot::PerPacketInfo = ppi.reader(mem).read_plain()?;
2655 if h.size == 0 {
2656 return Err(WorkerError::RndisMessageTooSmall(
2657 MessageComponentError::PerPacketInfo,
2658 ));
2659 }
2660 let (this, rest) =
2661 ppi.try_split(h.size as usize)
2662 .ok_or(WorkerError::RndisMessageTooSmall(
2663 MessageComponentError::PerPacketInfo,
2664 ))?;
2665 let (_, d) = this
2666 .try_split(h.per_packet_information_offset as usize)
2667 .ok_or(WorkerError::RndisMessageTooSmall(
2668 MessageComponentError::PerPacketInfo,
2669 ))?;
2670 match h.typ {
2671 rndisprot::PPI_TCP_IP_CHECKSUM => {
2672 let n: rndisprot::TxTcpIpChecksumInfo = d.reader(mem).read_plain()?;
2673
2674 metadata.flags.set_offload_tcp_checksum(
2675 (n.is_ipv4() || n.is_ipv6()) && n.tcp_checksum(),
2676 );
2677 metadata.flags.set_offload_udp_checksum(
2678 (n.is_ipv4() || n.is_ipv6()) && !n.tcp_checksum() && n.udp_checksum(),
2679 );
2680 metadata
2681 .flags
2682 .set_offload_ip_header_checksum(n.is_ipv4() && n.ip_header_checksum());
2683 metadata.flags.set_is_ipv4(n.is_ipv4());
2684 metadata.flags.set_is_ipv6(n.is_ipv6() && !n.is_ipv4());
2685 metadata.transport_header_offset = n.tcp_header_offset();
2686 }
2687 rndisprot::PPI_LSO => {
2688 let n: rndisprot::TcpLsoInfo = d.reader(mem).read_plain()?;
2689
2690 metadata.flags.set_offload_tcp_segmentation(true);
2691 metadata.flags.set_offload_tcp_checksum(true);
2692 metadata.flags.set_offload_ip_header_checksum(n.is_ipv4());
2693 metadata.flags.set_is_ipv4(n.is_ipv4());
2694 metadata.flags.set_is_ipv6(n.is_ipv6() && !n.is_ipv4());
2695 metadata.max_segment_size = n.mss() as u16;
2696 metadata.transport_header_offset = n.tcp_header_offset();
2697 }
2698 rndisprot::PPI_VLAN => {
2699 let n: rndisprot::EthVlanInfo = d.reader(mem).read_plain()?;
2700
2701 metadata.vlan = Some(n.into());
2702 }
2703 _ => {}
2704 }
2705 ppi = rest;
2706 }
2707
2708 metadata.l2_len = net_backend::ETHERNET_HEADER_LEN as u8;
2714
2715 if metadata.flags.offload_tcp_checksum() || metadata.flags.offload_udp_checksum() {
2716 metadata.l3_len = if metadata.transport_header_offset == 0 {
2720 if metadata.flags.is_ipv4() {
2721 net_backend::IPV4_MIN_HEADER_LEN
2722 } else if metadata.flags.is_ipv6() {
2723 net_backend::IPV6_MIN_HEADER_LEN
2724 } else {
2725 unreachable!("this packet is neither v4 nor v6?");
2726 }
2727 } else if (metadata.transport_header_offset < metadata.l2_len as u16)
2728 || (metadata.flags.is_ipv4()
2729 && metadata.transport_header_offset
2730 < (metadata.l2_len as u16 + net_backend::IPV4_MIN_HEADER_LEN))
2731 || (metadata.flags.is_ipv6()
2732 && metadata.transport_header_offset
2733 < (metadata.l2_len as u16 + net_backend::IPV6_MIN_HEADER_LEN))
2734 || (metadata.transport_header_offset as u32 >= request.data_length)
2735 {
2736 return Err(WorkerError::InvalidTcpHeaderOffset(
2737 metadata.transport_header_offset,
2738 ));
2739 } else {
2740 metadata.transport_header_offset - metadata.l2_len as u16
2741 }
2742 }
2743
2744 if metadata.flags.offload_tcp_segmentation() {
2745 const TCP_DOFF_BYTE_OFFSET: u32 = 12;
2746 let tcp_hdr_doff_offset =
2747 u32::from(metadata.transport_header_offset) + TCP_DOFF_BYTE_OFFSET;
2748 if tcp_hdr_doff_offset >= request.data_length {
2750 return Err(WorkerError::InvalidTcpHeaderOffset(
2751 metadata.transport_header_offset,
2752 ));
2753 }
2754 metadata.l4_len = {
2755 let mut reader = data.clone().reader(mem);
2756 reader.skip(tcp_hdr_doff_offset as usize)?;
2757 let mut b = 0;
2758 reader.read(std::slice::from_mut(&mut b))?;
2759 (b >> 4) * 4
2760 };
2761
2762 if request.data_length >= rndisprot::LSO_MAX_OFFLOAD_SIZE {
2763 stats.tx_invalid_lso_packets.increment();
2765 }
2766 }
2767
2768 }
2770
2771 let start = segments.len();
2772 for range in data.paged_ranges().flat_map(|r| r.ranges()) {
2773 let range = range.map_err(WorkerError::InvalidGpadl)?;
2774 segments.push(TxSegment {
2775 ty: net_backend::TxSegmentType::Tail,
2776 gpa: range.start,
2777 len: range.len() as u32,
2778 });
2779 }
2780
2781 metadata.segment_count = (segments.len() - start) as u8;
2782
2783 stats.tx_packets.increment();
2784 if metadata.flags.offload_tcp_checksum() || metadata.flags.offload_udp_checksum() {
2785 stats.tx_checksum_packets.increment();
2786 }
2787 if metadata.flags.offload_tcp_segmentation() {
2788 stats.tx_lso_packets.increment();
2789 }
2790 if metadata.vlan.is_some() {
2791 stats.tx_vlan_packets.increment();
2792 }
2793
2794 segments[start].ty = net_backend::TxSegmentType::Head(metadata);
2795
2796 Ok(())
2797 }
2798
2799 fn guest_vf_is_available(
2807 &mut self,
2808 primary: &mut PrimaryChannelState,
2809 vfid: Option<u32>,
2810 version: Version,
2811 config: NdisConfig,
2812 ) -> Result<bool, WorkerError> {
2813 let (serial_number, available) = if let Some(vfid) = vfid {
2814 (self.adapter.get_guest_vf_serial_number(vfid), true)
2815 } else {
2816 (primary.advertised_vf_serial_number.unwrap_or(0), false)
2817 };
2818 if version >= Version::V4 && config.capabilities.sriov() {
2819 tracing::info!(available, serial_number, "sending VF association message");
2820 let message = {
2822 self.message(
2823 protocol::MESSAGE4_TYPE_SEND_VF_ASSOCIATION,
2824 protocol::Message4SendVfAssociation {
2825 vf_allocated: if available { 1 } else { 0 },
2826 serial_number,
2827 },
2828 )
2829 };
2830 self.queue
2831 .split()
2832 .1
2833 .batched()
2834 .try_write_aligned(
2835 VF_ASSOCIATION_TRANSACTION_ID,
2836 OutgoingPacketType::InBandWithCompletion,
2837 message.aligned_payload(),
2838 )
2839 .map_err(|err| match err {
2840 queue::TryWriteError::Full(len) => {
2841 tracing::error!(len, "failed to write vf association message");
2842 WorkerError::OutOfSpace
2843 }
2844 queue::TryWriteError::Queue(err) => WorkerError::Queue(err),
2845 })?;
2846
2847 if available {
2849 primary.advertised_vf_serial_number = Some(serial_number);
2850 } else {
2851 primary.advertised_vf_serial_number = None;
2852 }
2853 Ok(true)
2854 } else {
2855 tracing::info!(
2856 available,
2857 serial_number,
2858 major = version.major(),
2859 minor = version.minor(),
2860 sriov_capable = config.capabilities.sriov(),
2861 "Skipping NvspMessage4TypeSendVFAssociation message"
2862 );
2863 Ok(false)
2864 }
2865 }
2866
2867 fn guest_send_indirection_table(&mut self, version: Version, num_channels_opened: u32) {
2869 if version < Version::V5 {
2871 return;
2872 }
2873
2874 #[repr(C)]
2875 #[derive(IntoBytes, Immutable, KnownLayout)]
2876 struct SendIndirectionMsg {
2877 pub message: protocol::Message5SendIndirectionTable,
2878 pub send_indirection_table:
2879 [u32; VMS_SWITCH_RSS_MAX_SEND_INDIRECTION_TABLE_ENTRIES as usize],
2880 }
2881
2882 let send_indirection_table_offset = offset_of!(SendIndirectionMsg, send_indirection_table)
2884 + size_of::<protocol::MessageHeader>();
2885 let mut data = SendIndirectionMsg {
2886 message: protocol::Message5SendIndirectionTable {
2887 table_entry_count: VMS_SWITCH_RSS_MAX_SEND_INDIRECTION_TABLE_ENTRIES,
2888 table_offset: send_indirection_table_offset as u32,
2889 },
2890 send_indirection_table: Default::default(),
2891 };
2892
2893 for i in 0..data.send_indirection_table.len() {
2894 data.send_indirection_table[i] = i as u32 % num_channels_opened;
2895 }
2896
2897 let header = protocol::MessageHeader {
2898 message_type: protocol::MESSAGE5_TYPE_SEND_INDIRECTION_TABLE,
2899 };
2900 let result = self
2901 .queue
2902 .split()
2903 .1
2904 .try_write(&queue::OutgoingPacket {
2905 transaction_id: 0,
2906 packet_type: OutgoingPacketType::InBandNoCompletion,
2907 payload: &[header.as_bytes(), data.as_bytes()],
2908 })
2909 .map_err(|err| match err {
2910 queue::TryWriteError::Full(len) => {
2911 tracing::error!(len, "failed to write send indirection table message");
2912 WorkerError::OutOfSpace
2913 }
2914 queue::TryWriteError::Queue(err) => WorkerError::Queue(err),
2915 });
2916 if let Err(err) = result {
2917 tracing::error!(
2918 error = &err as &dyn std::error::Error,
2919 "Failed to notify guest about the send indirection table"
2920 );
2921 }
2922 }
2923
2924 fn guest_vf_data_path_switched_to_synthetic(&mut self) {
2927 let header = protocol::MessageHeader {
2928 message_type: protocol::MESSAGE4_TYPE_SWITCH_DATA_PATH,
2929 };
2930 let data = protocol::Message4SwitchDataPath {
2931 active_data_path: protocol::DataPath::SYNTHETIC.0,
2932 };
2933 let result = self
2934 .queue
2935 .split()
2936 .1
2937 .try_write(&queue::OutgoingPacket {
2938 transaction_id: SWITCH_DATA_PATH_TRANSACTION_ID,
2939 packet_type: OutgoingPacketType::InBandWithCompletion,
2940 payload: &[header.as_bytes(), data.as_bytes()],
2941 })
2942 .map_err(|err| match err {
2943 queue::TryWriteError::Full(len) => {
2944 tracing::error!(len, "failed to write switch data path message");
2945 WorkerError::OutOfSpace
2946 }
2947 queue::TryWriteError::Queue(err) => WorkerError::Queue(err),
2948 });
2949 if let Err(err) = result {
2950 tracing::error!(
2951 error = &err as &dyn std::error::Error,
2952 "Failed to notify guest that data path is now synthetic"
2953 );
2954 } else {
2955 tracing::info!("Switched data path to synthetic")
2956 }
2957 }
2958
2959 async fn handle_state_change(
2961 &mut self,
2962 primary: &mut PrimaryChannelState,
2963 buffers: &ChannelBuffers,
2964 ) -> Result<Option<CoordinatorMessage>, WorkerError> {
2965 if let PrimaryChannelGuestVfState::Available { vfid } = primary.guest_vf_state {
2971 if primary.rndis_state == RndisState::Operational {
2973 if self.guest_vf_is_available(
2974 primary,
2975 Some(vfid),
2976 buffers.version,
2977 buffers.ndis_config,
2978 )? {
2979 primary.guest_vf_state = PrimaryChannelGuestVfState::AvailableAdvertised;
2980 return Ok(Some(CoordinatorMessage::Update(
2981 CoordinatorMessageUpdateType {
2982 guest_vf_state: true,
2983 ..Default::default()
2984 },
2985 )));
2986 } else if let Some(true) = primary.is_data_path_switched {
2987 tracing::error!(
2988 "Data path switched, but current guest negotiation does not support VTL0 VF"
2989 );
2990 }
2991 }
2992 return Ok(None);
2993 }
2994 loop {
2995 primary.guest_vf_state = match primary.guest_vf_state {
2996 PrimaryChannelGuestVfState::UnavailableFromAvailable => {
2997 if primary.rndis_state == RndisState::Operational {
2999 self.guest_vf_is_available(
3000 primary,
3001 None,
3002 buffers.version,
3003 buffers.ndis_config,
3004 )?;
3005 }
3006 PrimaryChannelGuestVfState::Unavailable
3007 }
3008 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending {
3009 to_guest,
3010 id,
3011 } => {
3012 self.send_completion(id, None)?;
3014 if to_guest {
3015 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
3016 } else {
3017 PrimaryChannelGuestVfState::UnavailableFromAvailable
3018 }
3019 }
3020 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched => {
3021 self.guest_vf_data_path_switched_to_synthetic();
3023 PrimaryChannelGuestVfState::UnavailableFromAvailable
3024 }
3025 PrimaryChannelGuestVfState::DataPathSynthetic => {
3026 self.guest_vf_data_path_switched_to_synthetic();
3028 PrimaryChannelGuestVfState::Ready
3029 }
3030 PrimaryChannelGuestVfState::DataPathSwitchPending {
3031 to_guest,
3032 id,
3033 result,
3034 } => {
3035 let result = result.expect("DataPathSwitchPending should have been processed");
3036 self.send_completion(id, None)?;
3038
3039 match (to_guest, result) {
3040 (true, true) => PrimaryChannelGuestVfState::DataPathSwitched,
3042 (true, false) => {
3044 tracing::error!(
3045 "Failure switching to guest VF, remaining on synthetic"
3046 );
3047 PrimaryChannelGuestVfState::DataPathSynthetic
3048 }
3049 (false, true) => PrimaryChannelGuestVfState::Ready,
3051 (false, false) => {
3053 tracing::error!(
3054 "Failure when guest requested switch back to synthetic"
3055 );
3056 PrimaryChannelGuestVfState::DataPathSwitched
3057 }
3058 }
3059 }
3060 PrimaryChannelGuestVfState::Initializing
3061 | PrimaryChannelGuestVfState::Restoring(_) => {
3062 panic!("Invalid guest VF state: {}", primary.guest_vf_state)
3063 }
3064 _ => break,
3065 };
3066 }
3067 Ok(None)
3068 }
3069
3070 fn handle_rndis_control_message(
3073 &mut self,
3074 primary: &mut PrimaryChannelState,
3075 buffers: &ChannelBuffers,
3076 message_type: u32,
3077 mut reader: impl MemoryRead + Clone,
3078 id: u32,
3079 ) -> Result<(), WorkerError> {
3080 let mem = &buffers.mem;
3081 let buffer_range = &buffers.recv_buffer.range(id);
3082 match message_type {
3083 rndisprot::MESSAGE_TYPE_INITIALIZE_MSG => {
3084 if primary.rndis_state != RndisState::Initializing {
3085 return Err(WorkerError::InvalidRndisState);
3086 }
3087
3088 let request: rndisprot::InitializeRequest = reader.read_plain()?;
3089
3090 tracing::trace!(
3091 ?request,
3092 "handling control message MESSAGE_TYPE_INITIALIZE_MSG"
3093 );
3094
3095 primary.rndis_state = RndisState::Operational;
3096
3097 let mut writer = buffer_range.writer(mem);
3098 let message_length = write_rndis_message(
3099 &mut writer,
3100 rndisprot::MESSAGE_TYPE_INITIALIZE_CMPLT,
3101 0,
3102 &rndisprot::InitializeComplete {
3103 request_id: request.request_id,
3104 status: rndisprot::STATUS_SUCCESS,
3105 major_version: rndisprot::MAJOR_VERSION,
3106 minor_version: rndisprot::MINOR_VERSION,
3107 device_flags: rndisprot::DF_CONNECTIONLESS,
3108 medium: rndisprot::MEDIUM_802_3,
3109 max_packets_per_message: 8,
3110 max_transfer_size: 0xEFFFFFFF,
3111 packet_alignment_factor: 3,
3112 af_list_offset: 0,
3113 af_list_size: 0,
3114 },
3115 )?;
3116 self.send_rndis_control_message(buffers, id, message_length)?;
3117 if let PrimaryChannelGuestVfState::Available { vfid } = primary.guest_vf_state {
3118 if self.guest_vf_is_available(
3119 primary,
3120 Some(vfid),
3121 buffers.version,
3122 buffers.ndis_config,
3123 )? {
3124 primary.guest_vf_state = PrimaryChannelGuestVfState::AvailableAdvertised;
3136 self.send_coordinator_update_vf();
3137 } else if let Some(true) = primary.is_data_path_switched {
3138 tracing::error!(
3139 "Data path switched, but current guest negotiation does not support VTL0 VF"
3140 );
3141 }
3142 }
3143 }
3144 rndisprot::MESSAGE_TYPE_QUERY_MSG => {
3145 let request: rndisprot::QueryRequest = reader.read_plain()?;
3146
3147 tracing::trace!(?request, "handling control message MESSAGE_TYPE_QUERY_MSG");
3148
3149 let (header, body) = buffer_range
3150 .try_split(
3151 size_of::<rndisprot::MessageHeader>()
3152 + size_of::<rndisprot::QueryComplete>(),
3153 )
3154 .ok_or(WorkerError::RndisMessageTooSmall(
3155 MessageComponentError::Header,
3156 ))?;
3157 let (status, tx) = match self.adapter.handle_oid_query(
3158 buffers,
3159 primary,
3160 request.oid,
3161 body.writer(mem),
3162 ) {
3163 Ok(tx) => (rndisprot::STATUS_SUCCESS, tx),
3164 Err(err) => (err.as_status(), 0),
3165 };
3166
3167 let message_length = write_rndis_message(
3168 &mut header.writer(mem),
3169 rndisprot::MESSAGE_TYPE_QUERY_CMPLT,
3170 tx,
3171 &rndisprot::QueryComplete {
3172 request_id: request.request_id,
3173 status,
3174 information_buffer_offset: size_of::<rndisprot::QueryComplete>() as u32,
3175 information_buffer_length: tx as u32,
3176 },
3177 )?;
3178 self.send_rndis_control_message(buffers, id, message_length)?;
3179 }
3180 rndisprot::MESSAGE_TYPE_SET_MSG => {
3181 let request: rndisprot::SetRequest = reader.read_plain()?;
3182
3183 tracing::trace!(?request, "handling control message MESSAGE_TYPE_SET_MSG");
3184
3185 let status = match self.adapter.handle_oid_set(primary, request.oid, reader) {
3186 Ok((restart_endpoint, packet_filter)) => {
3187 if restart_endpoint {
3190 self.restart = Some(CoordinatorMessage::Restart { channel_idx: 0 });
3191 }
3192 if let Some(filter) = packet_filter {
3193 if self.packet_filter != filter {
3194 self.packet_filter = filter;
3195 self.send_coordinator_update_filter();
3196 }
3197 }
3198 rndisprot::STATUS_SUCCESS
3199 }
3200 Err(err) => {
3201 tracelimit::warn_ratelimited!(
3202 error = &err as &dyn std::error::Error,
3203 oid = ?request.oid,
3204 "oid set failure"
3205 );
3206 err.as_status()
3207 }
3208 };
3209
3210 let message_length = write_rndis_message(
3211 &mut buffer_range.writer(mem),
3212 rndisprot::MESSAGE_TYPE_SET_CMPLT,
3213 0,
3214 &rndisprot::SetComplete {
3215 request_id: request.request_id,
3216 status,
3217 },
3218 )?;
3219 self.send_rndis_control_message(buffers, id, message_length)?;
3220 }
3221 rndisprot::MESSAGE_TYPE_RESET_MSG => {
3222 return Err(WorkerError::RndisMessageTypeNotImplemented);
3223 }
3224 rndisprot::MESSAGE_TYPE_INDICATE_STATUS_MSG => {
3225 return Err(WorkerError::RndisMessageTypeNotImplemented);
3226 }
3227 rndisprot::MESSAGE_TYPE_KEEPALIVE_MSG => {
3228 let request: rndisprot::KeepaliveRequest = reader.read_plain()?;
3229
3230 tracing::trace!(
3231 ?request,
3232 "handling control message MESSAGE_TYPE_KEEPALIVE_MSG"
3233 );
3234
3235 let message_length = write_rndis_message(
3236 &mut buffer_range.writer(mem),
3237 rndisprot::MESSAGE_TYPE_KEEPALIVE_CMPLT,
3238 0,
3239 &rndisprot::KeepaliveComplete {
3240 request_id: request.request_id,
3241 status: rndisprot::STATUS_SUCCESS,
3242 },
3243 )?;
3244 self.send_rndis_control_message(buffers, id, message_length)?;
3245 }
3246 rndisprot::MESSAGE_TYPE_SET_EX_MSG => {
3247 return Err(WorkerError::RndisMessageTypeNotImplemented);
3248 }
3249 _ => return Err(WorkerError::UnknownRndisMessageType(message_type)),
3250 };
3251 Ok(())
3252 }
3253
3254 fn try_send_rndis_message(
3255 &mut self,
3256 transaction_id: u64,
3257 channel_type: u32,
3258 recv_buffer_id: u16,
3259 transfer_pages: &[ring::TransferPageRange],
3260 ) -> Result<Option<usize>, WorkerError> {
3261 let message = self.message(
3262 protocol::MESSAGE1_TYPE_SEND_RNDIS_PACKET,
3263 protocol::Message1SendRndisPacket {
3264 channel_type,
3265 send_buffer_section_index: 0xffffffff,
3266 send_buffer_section_size: 0,
3267 },
3268 );
3269 let pending_send_size = match self.queue.split().1.batched().try_write_aligned(
3270 transaction_id,
3271 OutgoingPacketType::TransferPages(recv_buffer_id, transfer_pages),
3272 message.aligned_payload(),
3273 ) {
3274 Ok(()) => None,
3275 Err(queue::TryWriteError::Full(n)) => Some(n),
3276 Err(queue::TryWriteError::Queue(err)) => return Err(err.into()),
3277 };
3278 Ok(pending_send_size)
3279 }
3280
3281 fn send_rndis_control_message(
3282 &mut self,
3283 buffers: &ChannelBuffers,
3284 id: u32,
3285 message_length: usize,
3286 ) -> Result<(), WorkerError> {
3287 let result = self.try_send_rndis_message(
3288 id as u64,
3289 protocol::CONTROL_CHANNEL_TYPE,
3290 buffers.recv_buffer.id,
3291 std::slice::from_ref(&buffers.recv_buffer.transfer_page_range(id, message_length)),
3292 )?;
3293
3294 match result {
3296 None => Ok(()),
3297 Some(len) => {
3298 tracelimit::error_ratelimited!(len, "failed to write control message completion");
3299 Err(WorkerError::OutOfSpace)
3300 }
3301 }
3302 }
3303
3304 fn indicate_status(
3305 &mut self,
3306 buffers: &ChannelBuffers,
3307 id: u32,
3308 status: u32,
3309 payload: &[u8],
3310 ) -> Result<(), WorkerError> {
3311 let buffer = &buffers.recv_buffer.range(id);
3312 let mut writer = buffer.writer(&buffers.mem);
3313 let message_length = write_rndis_message(
3314 &mut writer,
3315 rndisprot::MESSAGE_TYPE_INDICATE_STATUS_MSG,
3316 payload.len(),
3317 &rndisprot::IndicateStatus {
3318 status,
3319 status_buffer_length: payload.len() as u32,
3320 status_buffer_offset: if payload.is_empty() {
3321 0
3322 } else {
3323 size_of::<rndisprot::IndicateStatus>() as u32
3324 },
3325 },
3326 )?;
3327 writer.write(payload)?;
3328 self.send_rndis_control_message(buffers, id, message_length)?;
3329 Ok(())
3330 }
3331
3332 fn process_control_messages(
3335 &mut self,
3336 buffers: &ChannelBuffers,
3337 state: &mut ActiveState,
3338 ) -> Result<(), WorkerError> {
3339 let Some(primary) = &mut state.primary else {
3340 return Ok(());
3341 };
3342
3343 while !primary.control_messages.is_empty()
3344 || (primary.pending_offload_change && primary.rndis_state == RndisState::Operational)
3345 {
3346 if !self.queue.split().1.can_write(MIN_CONTROL_RING_SIZE)? {
3348 self.pending_send_size = MIN_CONTROL_RING_SIZE;
3349 break;
3350 }
3351 let Some(id) = primary.free_control_buffers.pop() else {
3352 break;
3353 };
3354
3355 assert!(state.rx_bufs.is_free(id.0));
3357 state.rx_bufs.allocate(std::iter::once(id.0)).unwrap();
3358
3359 if let Some(message) = primary.control_messages.pop_front() {
3360 primary.control_messages_len -= message.data.len();
3361 self.handle_rndis_control_message(
3362 primary,
3363 buffers,
3364 message.message_type,
3365 message.data.as_ref(),
3366 id.0,
3367 )?;
3368 } else if primary.pending_offload_change
3369 && primary.rndis_state == RndisState::Operational
3370 {
3371 let ndis_offload = primary.offload_config.ndis_offload();
3372 self.indicate_status(
3373 buffers,
3374 id.0,
3375 rndisprot::STATUS_TASK_OFFLOAD_CURRENT_CONFIG,
3376 &ndis_offload.as_bytes()[..ndis_offload.header.size.into()],
3377 )?;
3378 primary.pending_offload_change = false;
3379 } else {
3380 unreachable!();
3381 }
3382 }
3383 Ok(())
3384 }
3385
3386 fn send_coordinator_update_message(&mut self, guest_vf: bool, packet_filter: bool) {
3387 if self.restart.is_none() {
3388 self.restart = Some(CoordinatorMessage::Update(CoordinatorMessageUpdateType {
3389 guest_vf_state: guest_vf,
3390 filter_state: packet_filter,
3391 }));
3392 } else if let Some(CoordinatorMessage::Restart { .. }) = self.restart {
3393 } else if let Some(CoordinatorMessage::Update(ref mut update)) = self.restart {
3397 update.guest_vf_state |= guest_vf;
3399 update.filter_state |= packet_filter;
3400 }
3401 }
3402
3403 fn send_coordinator_update_vf(&mut self) {
3404 self.send_coordinator_update_message(true, false);
3405 }
3406
3407 fn send_coordinator_update_filter(&mut self) {
3408 self.send_coordinator_update_message(false, true);
3409 }
3410}
3411
3412fn write_rndis_message<T: IntoBytes + Immutable + KnownLayout>(
3414 writer: &mut impl MemoryWrite,
3415 message_type: u32,
3416 extra: usize,
3417 payload: &T,
3418) -> Result<usize, AccessError> {
3419 let message_length = size_of::<rndisprot::MessageHeader>() + size_of_val(payload) + extra;
3420 writer.write(
3421 rndisprot::MessageHeader {
3422 message_type,
3423 message_length: message_length as u32,
3424 }
3425 .as_bytes(),
3426 )?;
3427 writer.write(payload.as_bytes())?;
3428 Ok(message_length)
3429}
3430
3431#[derive(Debug, Error)]
3432enum OidError {
3433 #[error(transparent)]
3434 Access(#[from] AccessError),
3435 #[error("unknown oid")]
3436 UnknownOid,
3437 #[error("invalid oid input, bad field {0}")]
3438 InvalidInput(&'static str),
3439 #[error("bad ndis version")]
3440 BadVersion,
3441 #[error("feature {0} not supported")]
3442 NotSupported(&'static str),
3443}
3444
3445impl OidError {
3446 fn as_status(&self) -> u32 {
3447 match self {
3448 OidError::UnknownOid | OidError::NotSupported(_) => rndisprot::STATUS_NOT_SUPPORTED,
3449 OidError::BadVersion => rndisprot::STATUS_BAD_VERSION,
3450 OidError::InvalidInput(_) => rndisprot::STATUS_INVALID_DATA,
3451 OidError::Access(_) => rndisprot::STATUS_FAILURE,
3452 }
3453 }
3454}
3455
3456const DEFAULT_MTU: u32 = 1514;
3457const MIN_MTU: u32 = DEFAULT_MTU;
3458const MAX_MTU: u32 = 9216;
3459
3460impl Adapter {
3461 fn get_guest_vf_serial_number(&self, vfid: u32) -> u32 {
3462 if let Some(guest_os_id) = self.get_guest_os_id.as_ref().map(|f| f()) {
3463 if guest_os_id
3466 .microsoft()
3467 .unwrap_or(HvGuestOsMicrosoft::from(0))
3468 .os_id()
3469 == HvGuestOsMicrosoftIds::WINDOWS_NT.0
3470 {
3471 self.adapter_index
3472 } else {
3473 vfid
3474 }
3475 } else {
3476 vfid
3477 }
3478 }
3479
3480 fn handle_oid_query(
3481 &self,
3482 buffers: &ChannelBuffers,
3483 primary: &PrimaryChannelState,
3484 oid: rndisprot::Oid,
3485 mut writer: impl MemoryWrite,
3486 ) -> Result<usize, OidError> {
3487 tracing::debug!(?oid, "oid query");
3488 let available_len = writer.len();
3489 match oid {
3490 rndisprot::Oid::OID_GEN_SUPPORTED_LIST => {
3491 let supported_oids_common = &[
3492 rndisprot::Oid::OID_GEN_SUPPORTED_LIST,
3493 rndisprot::Oid::OID_GEN_HARDWARE_STATUS,
3494 rndisprot::Oid::OID_GEN_MEDIA_SUPPORTED,
3495 rndisprot::Oid::OID_GEN_MEDIA_IN_USE,
3496 rndisprot::Oid::OID_GEN_MAXIMUM_LOOKAHEAD,
3497 rndisprot::Oid::OID_GEN_CURRENT_LOOKAHEAD,
3498 rndisprot::Oid::OID_GEN_MAXIMUM_FRAME_SIZE,
3499 rndisprot::Oid::OID_GEN_MAXIMUM_TOTAL_SIZE,
3500 rndisprot::Oid::OID_GEN_TRANSMIT_BLOCK_SIZE,
3501 rndisprot::Oid::OID_GEN_RECEIVE_BLOCK_SIZE,
3502 rndisprot::Oid::OID_GEN_LINK_SPEED,
3503 rndisprot::Oid::OID_GEN_TRANSMIT_BUFFER_SPACE,
3504 rndisprot::Oid::OID_GEN_RECEIVE_BUFFER_SPACE,
3505 rndisprot::Oid::OID_GEN_VENDOR_ID,
3506 rndisprot::Oid::OID_GEN_VENDOR_DESCRIPTION,
3507 rndisprot::Oid::OID_GEN_VENDOR_DRIVER_VERSION,
3508 rndisprot::Oid::OID_GEN_DRIVER_VERSION,
3509 rndisprot::Oid::OID_GEN_CURRENT_PACKET_FILTER,
3510 rndisprot::Oid::OID_GEN_PROTOCOL_OPTIONS,
3511 rndisprot::Oid::OID_GEN_MAC_OPTIONS,
3512 rndisprot::Oid::OID_GEN_MEDIA_CONNECT_STATUS,
3513 rndisprot::Oid::OID_GEN_MAXIMUM_SEND_PACKETS,
3514 rndisprot::Oid::OID_GEN_NETWORK_LAYER_ADDRESSES,
3515 rndisprot::Oid::OID_GEN_FRIENDLY_NAME,
3516 rndisprot::Oid::OID_802_3_PERMANENT_ADDRESS,
3518 rndisprot::Oid::OID_802_3_CURRENT_ADDRESS,
3519 rndisprot::Oid::OID_802_3_MULTICAST_LIST,
3520 rndisprot::Oid::OID_802_3_MAXIMUM_LIST_SIZE,
3521 rndisprot::Oid::OID_802_3_RCV_ERROR_ALIGNMENT,
3523 rndisprot::Oid::OID_802_3_XMIT_ONE_COLLISION,
3524 rndisprot::Oid::OID_802_3_XMIT_MORE_COLLISIONS,
3525 rndisprot::Oid::OID_GEN_RNDIS_CONFIG_PARAMETER,
3530 ];
3531
3532 let supported_oids_6 = &[
3534 rndisprot::Oid::OID_GEN_LINK_PARAMETERS,
3536 rndisprot::Oid::OID_GEN_LINK_STATE,
3537 rndisprot::Oid::OID_GEN_MAX_LINK_SPEED,
3538 rndisprot::Oid::OID_GEN_BYTES_RCV,
3540 rndisprot::Oid::OID_GEN_BYTES_XMIT,
3541 rndisprot::Oid::OID_TCP_OFFLOAD_PARAMETERS,
3543 rndisprot::Oid::OID_OFFLOAD_ENCAPSULATION,
3544 rndisprot::Oid::OID_TCP_OFFLOAD_HARDWARE_CAPABILITIES,
3545 rndisprot::Oid::OID_TCP_OFFLOAD_CURRENT_CONFIG,
3546 ];
3549
3550 let supported_oids_63 = &[
3551 rndisprot::Oid::OID_GEN_RECEIVE_SCALE_CAPABILITIES,
3552 rndisprot::Oid::OID_GEN_RECEIVE_SCALE_PARAMETERS,
3553 ];
3554
3555 match buffers.ndis_version.major {
3556 5 => {
3557 writer.write(supported_oids_common.as_bytes())?;
3558 }
3559 6 => {
3560 writer.write(supported_oids_common.as_bytes())?;
3561 writer.write(supported_oids_6.as_bytes())?;
3562 if buffers.ndis_version.minor >= 30 {
3563 writer.write(supported_oids_63.as_bytes())?;
3564 }
3565 }
3566 _ => return Err(OidError::BadVersion),
3567 }
3568 }
3569 rndisprot::Oid::OID_GEN_HARDWARE_STATUS => {
3570 let status: u32 = 0; writer.write(status.as_bytes())?;
3572 }
3573 rndisprot::Oid::OID_GEN_MEDIA_SUPPORTED | rndisprot::Oid::OID_GEN_MEDIA_IN_USE => {
3574 writer.write(rndisprot::MEDIUM_802_3.as_bytes())?;
3575 }
3576 rndisprot::Oid::OID_GEN_MAXIMUM_LOOKAHEAD
3577 | rndisprot::Oid::OID_GEN_CURRENT_LOOKAHEAD
3578 | rndisprot::Oid::OID_GEN_MAXIMUM_FRAME_SIZE => {
3579 let len: u32 = buffers.ndis_config.mtu - net_backend::ETHERNET_HEADER_LEN;
3580 writer.write(len.as_bytes())?;
3581 }
3582 rndisprot::Oid::OID_GEN_MAXIMUM_TOTAL_SIZE
3583 | rndisprot::Oid::OID_GEN_TRANSMIT_BLOCK_SIZE
3584 | rndisprot::Oid::OID_GEN_RECEIVE_BLOCK_SIZE => {
3585 let len: u32 = buffers.ndis_config.mtu;
3586 writer.write(len.as_bytes())?;
3587 }
3588 rndisprot::Oid::OID_GEN_LINK_SPEED => {
3589 let speed: u32 = (self.link_speed / 100) as u32; writer.write(speed.as_bytes())?;
3591 }
3592 rndisprot::Oid::OID_GEN_TRANSMIT_BUFFER_SPACE
3593 | rndisprot::Oid::OID_GEN_RECEIVE_BUFFER_SPACE => {
3594 writer.write((256u32 * 1024).as_bytes())?
3596 }
3597 rndisprot::Oid::OID_GEN_VENDOR_ID => {
3598 writer.write(0x0000155du32.as_bytes())?;
3601 }
3602 rndisprot::Oid::OID_GEN_VENDOR_DESCRIPTION => writer.write(b"Microsoft Corporation")?,
3603 rndisprot::Oid::OID_GEN_VENDOR_DRIVER_VERSION
3604 | rndisprot::Oid::OID_GEN_DRIVER_VERSION => {
3605 writer.write(0x0100u16.as_bytes())? }
3607 rndisprot::Oid::OID_GEN_CURRENT_PACKET_FILTER => writer.write(0u32.as_bytes())?,
3608 rndisprot::Oid::OID_GEN_MAC_OPTIONS => {
3609 let options: u32 = rndisprot::MAC_OPTION_COPY_LOOKAHEAD_DATA
3610 | rndisprot::MAC_OPTION_TRANSFERS_NOT_PEND
3611 | rndisprot::MAC_OPTION_NO_LOOPBACK
3612 | rndisprot::MAC_OPTION_8021P_PRIORITY
3613 | rndisprot::MAC_OPTION_8021Q_VLAN;
3614 writer.write(options.as_bytes())?;
3615 }
3616 rndisprot::Oid::OID_GEN_MEDIA_CONNECT_STATUS => {
3617 writer.write(rndisprot::MEDIA_STATE_CONNECTED.as_bytes())?;
3618 }
3619 rndisprot::Oid::OID_GEN_MAXIMUM_SEND_PACKETS => writer.write(u32::MAX.as_bytes())?,
3620 rndisprot::Oid::OID_GEN_FRIENDLY_NAME => {
3621 let name16: Vec<u16> = "Network Device".encode_utf16().collect();
3622 let mut name = rndisprot::FriendlyName::new_zeroed();
3623 name.name[..name16.len()].copy_from_slice(&name16);
3624 writer.write(name.as_bytes())?
3625 }
3626 rndisprot::Oid::OID_802_3_PERMANENT_ADDRESS
3627 | rndisprot::Oid::OID_802_3_CURRENT_ADDRESS => {
3628 writer.write(&self.mac_address.to_bytes())?
3629 }
3630 rndisprot::Oid::OID_802_3_MAXIMUM_LIST_SIZE => {
3631 writer.write(0u32.as_bytes())?;
3632 }
3633 rndisprot::Oid::OID_802_3_RCV_ERROR_ALIGNMENT
3634 | rndisprot::Oid::OID_802_3_XMIT_ONE_COLLISION
3635 | rndisprot::Oid::OID_802_3_XMIT_MORE_COLLISIONS => writer.write(0u32.as_bytes())?,
3636
3637 rndisprot::Oid::OID_GEN_LINK_STATE => {
3639 let link_state = rndisprot::LinkState {
3640 header: rndisprot::NdisObjectHeader {
3641 object_type: rndisprot::NdisObjectType::DEFAULT,
3642 revision: 1,
3643 size: size_of::<rndisprot::LinkState>() as u16,
3644 },
3645 media_connect_state: 1, media_duplex_state: 0, padding: 0,
3648 xmit_link_speed: self.link_speed,
3649 rcv_link_speed: self.link_speed,
3650 pause_functions: 0, auto_negotiation_flags: 0,
3652 };
3653 writer.write(link_state.as_bytes())?;
3654 }
3655 rndisprot::Oid::OID_GEN_MAX_LINK_SPEED => {
3656 let link_speed = rndisprot::LinkSpeed {
3657 xmit: self.link_speed,
3658 rcv: self.link_speed,
3659 };
3660 writer.write(link_speed.as_bytes())?;
3661 }
3662 rndisprot::Oid::OID_TCP_OFFLOAD_HARDWARE_CAPABILITIES => {
3663 let ndis_offload = self.offload_support.ndis_offload();
3664 writer.write(&ndis_offload.as_bytes()[..ndis_offload.header.size.into()])?;
3665 }
3666 rndisprot::Oid::OID_TCP_OFFLOAD_CURRENT_CONFIG => {
3667 let ndis_offload = primary.offload_config.ndis_offload();
3668 writer.write(&ndis_offload.as_bytes()[..ndis_offload.header.size.into()])?;
3669 }
3670 rndisprot::Oid::OID_OFFLOAD_ENCAPSULATION => {
3671 writer.write(
3672 &rndisprot::NdisOffloadEncapsulation {
3673 header: rndisprot::NdisObjectHeader {
3674 object_type: rndisprot::NdisObjectType::OFFLOAD_ENCAPSULATION,
3675 revision: 1,
3676 size: rndisprot::NDIS_SIZEOF_OFFLOAD_ENCAPSULATION_REVISION_1 as u16,
3677 },
3678 ipv4_enabled: rndisprot::NDIS_OFFLOAD_SUPPORTED,
3679 ipv4_encapsulation_type: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
3680 ipv4_header_size: net_backend::ETHERNET_HEADER_LEN,
3681 ipv6_enabled: rndisprot::NDIS_OFFLOAD_SUPPORTED,
3682 ipv6_encapsulation_type: rndisprot::NDIS_ENCAPSULATION_IEEE_802_3,
3683 ipv6_header_size: net_backend::ETHERNET_HEADER_LEN,
3684 }
3685 .as_bytes()[..rndisprot::NDIS_SIZEOF_OFFLOAD_ENCAPSULATION_REVISION_1],
3686 )?;
3687 }
3688 rndisprot::Oid::OID_GEN_RECEIVE_SCALE_CAPABILITIES => {
3689 writer.write(
3690 &rndisprot::NdisReceiveScaleCapabilities {
3691 header: rndisprot::NdisObjectHeader {
3692 object_type: rndisprot::NdisObjectType::RSS_CAPABILITIES,
3693 revision: 2,
3694 size: rndisprot::NDIS_SIZEOF_RECEIVE_SCALE_CAPABILITIES_REVISION_2
3695 as u16,
3696 },
3697 capabilities_flags: rndisprot::NDIS_RSS_CAPS_HASH_TYPE_TCP_IPV4
3698 | rndisprot::NDIS_RSS_CAPS_HASH_TYPE_TCP_IPV6
3699 | rndisprot::NDIS_HASH_FUNCTION_TOEPLITZ,
3700 number_of_interrupt_messages: 1,
3701 number_of_receive_queues: self.max_queues.into(),
3702 number_of_indirection_table_entries: if self.indirection_table_size != 0 {
3703 self.indirection_table_size
3704 } else {
3705 128
3708 },
3709 padding: 0,
3710 }
3711 .as_bytes()[..rndisprot::NDIS_SIZEOF_RECEIVE_SCALE_CAPABILITIES_REVISION_2],
3712 )?;
3713 }
3714 _ => {
3715 tracelimit::warn_ratelimited!(?oid, "query for unknown OID");
3716 return Err(OidError::UnknownOid);
3717 }
3718 };
3719 Ok(available_len - writer.len())
3720 }
3721
3722 fn handle_oid_set(
3723 &self,
3724 primary: &mut PrimaryChannelState,
3725 oid: rndisprot::Oid,
3726 reader: impl MemoryRead + Clone,
3727 ) -> Result<(bool, Option<u32>), OidError> {
3728 tracing::debug!(?oid, "oid set");
3729
3730 let mut restart_endpoint = false;
3731 let mut packet_filter = None;
3732 match oid {
3733 rndisprot::Oid::OID_GEN_CURRENT_PACKET_FILTER => {
3734 packet_filter = self.oid_set_packet_filter(reader)?;
3735 }
3736 rndisprot::Oid::OID_TCP_OFFLOAD_PARAMETERS => {
3737 self.oid_set_offload_parameters(reader, primary)?;
3738 }
3739 rndisprot::Oid::OID_OFFLOAD_ENCAPSULATION => {
3740 self.oid_set_offload_encapsulation(reader)?;
3741 }
3742 rndisprot::Oid::OID_GEN_RNDIS_CONFIG_PARAMETER => {
3743 self.oid_set_rndis_config_parameter(reader, primary)?;
3744 }
3745 rndisprot::Oid::OID_GEN_NETWORK_LAYER_ADDRESSES => {
3746 }
3748 rndisprot::Oid::OID_GEN_RECEIVE_SCALE_PARAMETERS => {
3749 let rss_was_enabled = self.oid_set_rss_parameters(reader, primary)?;
3750
3751 if rss_was_enabled || primary.rss_state.is_some() {
3758 restart_endpoint = true;
3759 }
3760 }
3761 _ => {
3762 tracelimit::warn_ratelimited!(?oid, "set of unknown OID");
3763 return Err(OidError::UnknownOid);
3764 }
3765 }
3766 Ok((restart_endpoint, packet_filter))
3767 }
3768
3769 fn oid_set_rss_parameters(
3770 &self,
3771 mut reader: impl MemoryRead + Clone,
3772 primary: &mut PrimaryChannelState,
3773 ) -> Result<bool, OidError> {
3774 let mut params = rndisprot::NdisReceiveScaleParameters::new_zeroed();
3776 let len = reader.len().min(size_of_val(¶ms));
3777 reader.clone().read(&mut params.as_mut_bytes()[..len])?;
3778
3779 let rss_was_enabled = primary.rss_state.is_some();
3780
3781 if ((params.flags & NDIS_RSS_PARAM_FLAG_DISABLE_RSS) != 0)
3782 || ((params.hash_information & NDIS_HASH_FUNCTION_MASK) == 0)
3783 {
3784 primary.rss_state = None;
3785 return Ok(rss_was_enabled);
3786 }
3787
3788 if params.hash_secret_key_size != 40 {
3789 return Err(OidError::InvalidInput("hash_secret_key_size"));
3790 }
3791 if params.indirection_table_size % 4 != 0 || params.indirection_table_size == 0 {
3792 return Err(OidError::InvalidInput("indirection_table_size"));
3793 }
3794 let indirection_table_size =
3795 (params.indirection_table_size / 4).min(self.indirection_table_size) as usize;
3796 let mut key = [0; 40];
3797 let mut indirection_table = vec![0u32; self.indirection_table_size as usize];
3798 reader
3799 .clone()
3800 .skip(params.hash_secret_key_offset as usize)?
3801 .read(&mut key)?;
3802 reader
3803 .skip(params.indirection_table_offset as usize)?
3804 .read(indirection_table[..indirection_table_size].as_mut_bytes())?;
3805 tracelimit::info_ratelimited!(?indirection_table, "OID_GEN_RECEIVE_SCALE_PARAMETERS");
3806 if indirection_table
3807 .iter()
3808 .any(|&x| x >= self.max_queues as u32)
3809 {
3810 return Err(OidError::InvalidInput("indirection_table"));
3811 }
3812 let (indir_init, indir_uninit) = indirection_table.split_at_mut(indirection_table_size);
3813 for (src, dest) in std::iter::repeat_with(|| indir_init.iter().copied())
3814 .flatten()
3815 .zip(indir_uninit)
3816 {
3817 *dest = src;
3818 }
3819 primary.rss_state = Some(RssState {
3820 key,
3821 indirection_table: indirection_table.iter().map(|&x| x as u16).collect(),
3822 });
3823 Ok(rss_was_enabled)
3824 }
3825
3826 fn oid_set_packet_filter(
3827 &self,
3828 reader: impl MemoryRead + Clone,
3829 ) -> Result<Option<u32>, OidError> {
3830 let filter: rndisprot::RndisPacketFilterOidValue = reader.clone().read_plain()?;
3831 tracing::debug!(filter, "set packet filter");
3832 Ok(Some(filter))
3833 }
3834
3835 fn oid_set_offload_parameters(
3836 &self,
3837 reader: impl MemoryRead + Clone,
3838 primary: &mut PrimaryChannelState,
3839 ) -> Result<(), OidError> {
3840 let offload: rndisprot::NdisOffloadParameters = read_ndis_object(
3841 reader,
3842 rndisprot::NdisObjectType::DEFAULT,
3843 1,
3844 rndisprot::NDIS_SIZEOF_OFFLOAD_PARAMETERS_REVISION_1,
3845 )?;
3846
3847 tracing::debug!(?offload, "offload parameters");
3848 let rndisprot::NdisOffloadParameters {
3849 header: _,
3850 ipv4_checksum,
3851 tcp4_checksum,
3852 udp4_checksum,
3853 tcp6_checksum,
3854 udp6_checksum,
3855 lsov1,
3856 ipsec_v1: _,
3857 lsov2_ipv4,
3858 lsov2_ipv6,
3859 tcp_connection_ipv4: _,
3860 tcp_connection_ipv6: _,
3861 reserved: _,
3862 flags: _,
3863 } = offload;
3864
3865 if lsov1 == rndisprot::OffloadParametersSimple::ENABLED {
3866 return Err(OidError::NotSupported("lsov1"));
3867 }
3868 if let Some((tx, rx)) = ipv4_checksum.tx_rx() {
3869 primary.offload_config.checksum_tx.ipv4_header = tx;
3870 primary.offload_config.checksum_rx.ipv4_header = rx;
3871 }
3872 if let Some((tx, rx)) = tcp4_checksum.tx_rx() {
3873 primary.offload_config.checksum_tx.tcp4 = tx;
3874 primary.offload_config.checksum_rx.tcp4 = rx;
3875 }
3876 if let Some((tx, rx)) = tcp6_checksum.tx_rx() {
3877 primary.offload_config.checksum_tx.tcp6 = tx;
3878 primary.offload_config.checksum_rx.tcp6 = rx;
3879 }
3880 if let Some((tx, rx)) = udp4_checksum.tx_rx() {
3881 primary.offload_config.checksum_tx.udp4 = tx;
3882 primary.offload_config.checksum_rx.udp4 = rx;
3883 }
3884 if let Some((tx, rx)) = udp6_checksum.tx_rx() {
3885 primary.offload_config.checksum_tx.udp6 = tx;
3886 primary.offload_config.checksum_rx.udp6 = rx;
3887 }
3888 if let Some(enable) = lsov2_ipv4.enable() {
3889 primary.offload_config.lso4 = enable;
3890 }
3891 if let Some(enable) = lsov2_ipv6.enable() {
3892 primary.offload_config.lso6 = enable;
3893 }
3894 primary
3895 .offload_config
3896 .mask_to_supported(&self.offload_support);
3897 primary.pending_offload_change = true;
3898 Ok(())
3899 }
3900
3901 fn oid_set_offload_encapsulation(
3902 &self,
3903 reader: impl MemoryRead + Clone,
3904 ) -> Result<(), OidError> {
3905 let encap: rndisprot::NdisOffloadEncapsulation = read_ndis_object(
3906 reader,
3907 rndisprot::NdisObjectType::OFFLOAD_ENCAPSULATION,
3908 1,
3909 rndisprot::NDIS_SIZEOF_OFFLOAD_ENCAPSULATION_REVISION_1,
3910 )?;
3911 if encap.ipv4_enabled == rndisprot::NDIS_OFFLOAD_SET_ON
3912 && (encap.ipv4_encapsulation_type != rndisprot::NDIS_ENCAPSULATION_IEEE_802_3
3913 || encap.ipv4_header_size != net_backend::ETHERNET_HEADER_LEN)
3914 {
3915 return Err(OidError::NotSupported("ipv4 encap"));
3916 }
3917 if encap.ipv6_enabled == rndisprot::NDIS_OFFLOAD_SET_ON
3918 && (encap.ipv6_encapsulation_type != rndisprot::NDIS_ENCAPSULATION_IEEE_802_3
3919 || encap.ipv6_header_size != net_backend::ETHERNET_HEADER_LEN)
3920 {
3921 return Err(OidError::NotSupported("ipv6 encap"));
3922 }
3923 Ok(())
3924 }
3925
3926 fn oid_set_rndis_config_parameter(
3927 &self,
3928 reader: impl MemoryRead + Clone,
3929 primary: &mut PrimaryChannelState,
3930 ) -> Result<(), OidError> {
3931 let info: rndisprot::RndisConfigParameterInfo = reader.clone().read_plain()?;
3932 if info.name_length > 255 {
3933 return Err(OidError::InvalidInput("name_length"));
3934 }
3935 if info.value_length > 255 {
3936 return Err(OidError::InvalidInput("value_length"));
3937 }
3938 let name = reader
3939 .clone()
3940 .skip(info.name_offset as usize)?
3941 .read_n::<u16>(info.name_length as usize / 2)?;
3942 let name = String::from_utf16(&name).map_err(|_| OidError::InvalidInput("name"))?;
3943 let mut value = reader;
3944 value.skip(info.value_offset as usize)?;
3945 let mut value = value.limit(info.value_length as usize);
3946 match info.parameter_type {
3947 rndisprot::NdisParameterType::STRING => {
3948 let value = value.read_n::<u16>(info.value_length as usize / 2)?;
3949 let value =
3950 String::from_utf16(&value).map_err(|_| OidError::InvalidInput("value"))?;
3951 let as_num = value
3952 .as_bytes()
3953 .first()
3954 .map(|c| c.wrapping_sub(b'0'))
3955 .filter(|&c| c <= 9)
3956 .ok_or(OidError::InvalidInput("value as num"))?;
3957 let tx = as_num & 1 != 0;
3958 let rx = as_num & 2 != 0;
3959
3960 tracing::debug!(name, value, "rndis config");
3961 match name.as_str() {
3962 "*IPChecksumOffloadIPv4" => {
3963 primary.offload_config.checksum_tx.ipv4_header = tx;
3964 primary.offload_config.checksum_rx.ipv4_header = rx;
3965 }
3966 "*LsoV2IPv4" => {
3967 primary.offload_config.lso4 = as_num != 0;
3968 }
3969 "*LsoV2IPv6" => {
3970 primary.offload_config.lso6 = as_num != 0;
3971 }
3972 "*TCPChecksumOffloadIPv4" => {
3973 primary.offload_config.checksum_tx.tcp4 = tx;
3974 primary.offload_config.checksum_rx.tcp4 = rx;
3975 }
3976 "*TCPChecksumOffloadIPv6" => {
3977 primary.offload_config.checksum_tx.tcp6 = tx;
3978 primary.offload_config.checksum_rx.tcp6 = rx;
3979 }
3980 "*UDPChecksumOffloadIPv4" => {
3981 primary.offload_config.checksum_tx.udp4 = tx;
3982 primary.offload_config.checksum_rx.udp4 = rx;
3983 }
3984 "*UDPChecksumOffloadIPv6" => {
3985 primary.offload_config.checksum_tx.udp6 = tx;
3986 primary.offload_config.checksum_rx.udp6 = rx;
3987 }
3988 _ => {}
3989 }
3990 primary
3991 .offload_config
3992 .mask_to_supported(&self.offload_support);
3993 }
3994 rndisprot::NdisParameterType::INTEGER => {
3995 let value: u32 = value.read_plain()?;
3996 tracing::debug!(name, value, "rndis config");
3997 }
3998 parameter_type => tracelimit::warn_ratelimited!(
3999 name,
4000 ?parameter_type,
4001 "unhandled rndis config parameter type"
4002 ),
4003 }
4004 Ok(())
4005 }
4006}
4007
4008fn read_ndis_object<T: IntoBytes + FromBytes + Debug + Immutable + KnownLayout>(
4009 mut reader: impl MemoryRead,
4010 object_type: rndisprot::NdisObjectType,
4011 min_revision: u8,
4012 min_size: usize,
4013) -> Result<T, OidError> {
4014 let mut buffer = T::new_zeroed();
4015 let sent_size = reader.len();
4016 let len = sent_size.min(size_of_val(&buffer));
4017 reader.read(&mut buffer.as_mut_bytes()[..len])?;
4018 validate_ndis_object_header(
4019 &rndisprot::NdisObjectHeader::read_from_prefix(buffer.as_bytes())
4020 .unwrap()
4021 .0, sent_size,
4023 object_type,
4024 min_revision,
4025 min_size,
4026 )?;
4027 Ok(buffer)
4028}
4029
4030fn validate_ndis_object_header(
4031 header: &rndisprot::NdisObjectHeader,
4032 sent_size: usize,
4033 object_type: rndisprot::NdisObjectType,
4034 min_revision: u8,
4035 min_size: usize,
4036) -> Result<(), OidError> {
4037 if header.object_type != object_type {
4038 return Err(OidError::InvalidInput("header.object_type"));
4039 }
4040 if sent_size < header.size as usize {
4041 return Err(OidError::InvalidInput("header.size"));
4042 }
4043 if header.revision < min_revision {
4044 return Err(OidError::InvalidInput("header.revision"));
4045 }
4046 if (header.size as usize) < min_size {
4047 return Err(OidError::InvalidInput("header.size"));
4048 }
4049 Ok(())
4050}
4051
4052struct Coordinator {
4053 recv: mpsc::Receiver<CoordinatorMessage>,
4054 channel_control: ChannelControl,
4055 restart: bool,
4056 workers: Vec<TaskControl<NetQueue, Worker<GpadlRingMem>>>,
4057 buffers: Option<Arc<ChannelBuffers>>,
4058 num_queues: u16,
4059 active_packet_filter: u32,
4060 sleep_deadline: Option<Instant>,
4061}
4062
4063enum CoordinatorStatePendingVfState {
4067 Ready,
4069 Delay {
4071 timer: PolledTimer,
4072 delay_until: Instant,
4073 },
4074 Pending,
4076}
4077
4078struct CoordinatorState {
4079 endpoint: Box<dyn Endpoint>,
4080 adapter: Arc<Adapter>,
4081 virtual_function: Option<Box<dyn VirtualFunction>>,
4082 pending_vf_state: CoordinatorStatePendingVfState,
4083}
4084
4085impl InspectTaskMut<Coordinator> for CoordinatorState {
4086 fn inspect_mut(
4087 &mut self,
4088 req: inspect::Request<'_>,
4089 mut coordinator: Option<&mut Coordinator>,
4090 ) {
4091 let mut resp = req.respond();
4092
4093 let adapter = self.adapter.as_ref();
4094 resp.field("mac_address", adapter.mac_address)
4095 .field("max_queues", adapter.max_queues)
4096 .sensitivity_field(
4097 "offload_support",
4098 SensitivityLevel::Safe,
4099 &adapter.offload_support,
4100 )
4101 .field_mut_with("ring_size_limit", |v| -> anyhow::Result<_> {
4102 if let Some(v) = v {
4103 let v: usize = v.parse()?;
4104 adapter.ring_size_limit.store(v, Ordering::Relaxed);
4105 if let Some(this) = &mut coordinator {
4107 for worker in &mut this.workers {
4108 worker.update_with(|_, _| ());
4109 }
4110 }
4111 }
4112 Ok(adapter.ring_size_limit.load(Ordering::Relaxed))
4113 });
4114
4115 resp.field("endpoint_type", self.endpoint.endpoint_type())
4116 .field(
4117 "endpoint_max_queues",
4118 self.endpoint.multiqueue_support().max_queues,
4119 )
4120 .sensitivity_field_mut("endpoint", SensitivityLevel::Safe, self.endpoint.as_mut());
4121
4122 if let Some(coordinator) = coordinator {
4123 resp.sensitivity_child("queues", SensitivityLevel::Safe, |req| {
4124 let mut resp = req.respond();
4125 for (i, q) in coordinator.workers[..coordinator.num_queues as usize]
4126 .iter_mut()
4127 .enumerate()
4128 {
4129 resp.field_mut(&i.to_string(), q);
4130 }
4131 });
4132
4133 resp.merge(inspect::adhoc_mut(|req| {
4135 let deferred = req.defer();
4136 coordinator.workers[0].update_with(|_, worker| {
4137 let Some(worker) = worker.as_deref() else {
4138 return;
4139 };
4140 if let Some(state) = worker.state.ready() {
4141 deferred.respond(|resp| {
4142 resp.merge(&state.buffers);
4143 resp.sensitivity_field(
4144 "primary_channel_state",
4145 SensitivityLevel::Safe,
4146 &state.state.primary,
4147 )
4148 .sensitivity_field(
4149 "packet_filter",
4150 SensitivityLevel::Safe,
4151 inspect::AsHex(worker.channel.packet_filter),
4152 );
4153 });
4154 }
4155 })
4156 }));
4157 }
4158 }
4159}
4160
4161impl AsyncRun<Coordinator> for CoordinatorState {
4162 async fn run(
4163 &mut self,
4164 stop: &mut StopTask<'_>,
4165 coordinator: &mut Coordinator,
4166 ) -> Result<(), task_control::Cancelled> {
4167 coordinator.process(stop, self).await
4168 }
4169}
4170
4171impl Coordinator {
4172 async fn process(
4173 &mut self,
4174 stop: &mut StopTask<'_>,
4175 state: &mut CoordinatorState,
4176 ) -> Result<(), task_control::Cancelled> {
4177 loop {
4178 if self.restart {
4183 self.restart_worker_queues(stop, state).await?;
4184 }
4185
4186 for worker in &mut self.workers[1..] {
4189 worker.start();
4190 }
4191 if !self.workers[0].is_running()
4192 && self.workers[0].state().is_none_or(|worker| {
4193 !matches!(worker.state, WorkerState::WaitingForCoordinator(_))
4194 })
4195 {
4196 self.workers[0].start();
4197 }
4198
4199 enum Message {
4200 Internal(CoordinatorMessage),
4201 ChannelDisconnected,
4202 UpdateFromEndpoint(EndpointAction),
4203 UpdateFromVf(Rpc<(), ()>),
4204 OfferVfDevice,
4205 PendingVfStateComplete,
4206 TimerExpired,
4207 }
4208 let message = if matches!(
4209 state.pending_vf_state,
4210 CoordinatorStatePendingVfState::Pending
4211 ) {
4212 state
4215 .virtual_function
4216 .as_mut()
4217 .expect("Pending requires a VF")
4218 .guest_ready_for_device()
4219 .await;
4220 Message::PendingVfStateComplete
4221 } else {
4222 let timer_sleep = async {
4223 if let Some(deadline) = self.sleep_deadline {
4224 let mut timer = PolledTimer::new(&state.adapter.driver);
4225 timer.sleep_until(deadline).await;
4226 } else {
4227 pending::<()>().await;
4228 }
4229 Message::TimerExpired
4230 };
4231 let wait_for_message = async {
4232 let internal_msg = self
4233 .recv
4234 .next()
4235 .map(|x| x.map_or(Message::ChannelDisconnected, Message::Internal));
4236 let endpoint_restart = state
4237 .endpoint
4238 .wait_for_endpoint_action()
4239 .map(Message::UpdateFromEndpoint);
4240 if let Some(vf) = state.virtual_function.as_mut() {
4241 match state.pending_vf_state {
4242 CoordinatorStatePendingVfState::Ready
4243 | CoordinatorStatePendingVfState::Delay { .. } => {
4244 let offer_device = async {
4245 if let CoordinatorStatePendingVfState::Delay {
4246 timer,
4247 delay_until,
4248 } = &mut state.pending_vf_state
4249 {
4250 timer.sleep_until(*delay_until).await;
4251 } else {
4252 pending::<()>().await;
4253 }
4254 Message::OfferVfDevice
4255 };
4256 (
4257 internal_msg,
4258 offer_device,
4259 endpoint_restart,
4260 vf.wait_for_state_change().map(Message::UpdateFromVf),
4261 timer_sleep,
4262 )
4263 .race()
4264 .await
4265 }
4266 CoordinatorStatePendingVfState::Pending => unreachable!(),
4267 }
4268 } else {
4269 (internal_msg, endpoint_restart, timer_sleep).race().await
4270 }
4271 };
4272
4273 stop.until_stopped(wait_for_message).await?
4274 };
4275 match message {
4276 Message::Internal(msg) => {
4277 self.handle_coordinator_message(msg, state).await;
4278 self.handle_queued_coordinator_messages(state).await;
4282 }
4283 Message::UpdateFromVf(rpc) => {
4284 rpc.handle(async |_| {
4285 self.update_guest_vf_state(state).await;
4286 })
4287 .await;
4288 }
4289 Message::OfferVfDevice => {
4290 self.stop_primary_worker().await;
4291 if let Some(primary) = self.primary_mut() {
4292 if matches!(
4293 primary.guest_vf_state,
4294 PrimaryChannelGuestVfState::AvailableAdvertised
4295 ) {
4296 primary.guest_vf_state = PrimaryChannelGuestVfState::Ready;
4297 }
4298 }
4299
4300 state.pending_vf_state = CoordinatorStatePendingVfState::Pending;
4301 }
4302 Message::PendingVfStateComplete => {
4303 state.pending_vf_state = CoordinatorStatePendingVfState::Ready;
4305 }
4306 Message::TimerExpired => {
4307 self.stop_primary_worker().await;
4309 if let Some(primary) = self.primary_mut() {
4310 if let PendingLinkAction::Delay(up) = primary.pending_link_action {
4311 primary.pending_link_action = PendingLinkAction::Active(up);
4312 }
4313 }
4314 self.sleep_deadline = None;
4315 }
4316 Message::UpdateFromEndpoint(endpoint_action) => {
4317 self.handle_endpoint_action(endpoint_action).await;
4318 }
4319 Message::ChannelDisconnected => {
4320 break;
4321 }
4322 };
4323 }
4324 Ok(())
4325 }
4326
4327 async fn handle_endpoint_action(&mut self, action: EndpointAction) {
4328 match action {
4329 EndpointAction::RestartRequired => self.restart = true,
4330 EndpointAction::LinkStatusNotify(connect) => {
4331 self.stop_primary_worker().await;
4332 if let Some(primary) = self.primary_mut() {
4340 primary.pending_link_action = PendingLinkAction::Active(connect);
4341 }
4342
4343 self.sleep_deadline = None;
4345 }
4346 }
4347 }
4348
4349 async fn restart_worker_queues(
4352 &mut self,
4353 stop: &mut StopTask<'_>,
4354 state: &mut CoordinatorState,
4355 ) -> Result<(), task_control::Cancelled> {
4356 stop.until_stopped(self.stop_workers()).await?;
4357
4358 self.handle_queued_coordinator_messages(state).await;
4362
4363 while let Some(action) = state.endpoint.wait_for_endpoint_action().now_or_never() {
4366 self.handle_endpoint_action(action).await;
4367 }
4368
4369 if let Err(err) = self
4371 .restart_queues(state)
4372 .instrument(tracing::info_span!("netvsp_restart_queues"))
4373 .await
4374 {
4375 tracing::error!(
4376 error = &err as &dyn std::error::Error,
4377 "failed to restart queues"
4378 );
4379 }
4380 if let Some(primary) = self.primary_mut() {
4381 primary.is_data_path_switched = state.endpoint.get_data_path_to_guest_vf().await.ok();
4382 tracing::info!(
4383 is_data_path_switched = primary.is_data_path_switched,
4384 "Query data path state"
4385 );
4386 }
4387 self.restore_guest_vf_state(state).await;
4388 self.restart = false;
4389 Ok(())
4390 }
4391
4392 async fn handle_queued_coordinator_messages(&mut self, state: &mut CoordinatorState) {
4393 while let Ok(Some(msg)) = self.recv.try_next() {
4394 self.handle_coordinator_message(msg, state).await;
4395 }
4396 }
4397
4398 async fn handle_coordinator_message(
4399 &mut self,
4400 msg: CoordinatorMessage,
4401 state: &mut CoordinatorState,
4402 ) {
4403 match msg {
4404 CoordinatorMessage::Restart { channel_idx } if channel_idx != 0 => {
4405 tracelimit::event_ratelimited!(
4406 tracing::Level::DEBUG,
4407 channel_idx,
4408 "sub-channel triggered restart"
4409 );
4410 self.restart = true;
4411 }
4412 _ => self.handle_primary_message(msg, state).await,
4413 }
4414 }
4415
4416 async fn handle_primary_message(
4417 &mut self,
4418 msg: CoordinatorMessage,
4419 state: &mut CoordinatorState,
4420 ) {
4421 self.stop_primary_worker().await;
4422 if let Some(worker) = self.workers[0].state_mut() {
4423 if matches!(worker.state, WorkerState::WaitingForCoordinator(_)) {
4424 let WorkerState::WaitingForCoordinator(Some(ready)) =
4425 std::mem::replace(&mut worker.state, WorkerState::WaitingForCoordinator(None))
4426 else {
4427 unreachable!("valid ready state")
4428 };
4429 let _ = std::mem::replace(&mut worker.state, WorkerState::Ready(ready));
4430 }
4431 }
4432 match msg {
4433 CoordinatorMessage::Update(update_type) => {
4434 if update_type.filter_state {
4435 self.stop_workers().await;
4436 self.active_packet_filter =
4437 self.workers[0].state().unwrap().channel.packet_filter;
4438 self.workers.iter_mut().skip(1).for_each(|worker| {
4439 if let Some(state) = worker.state_mut() {
4440 state.channel.packet_filter = self.active_packet_filter;
4441 tracing::debug!(
4442 packet_filter = ?self.active_packet_filter,
4443 channel_idx = state.channel_idx,
4444 "update packet filter"
4445 );
4446 }
4447 });
4448 }
4449
4450 if update_type.guest_vf_state {
4451 self.update_guest_vf_state(state).await;
4452 }
4453 }
4454 CoordinatorMessage::StartTimer(deadline) => {
4455 self.sleep_deadline = Some(deadline);
4456 }
4457 CoordinatorMessage::Restart { channel_idx } => {
4458 assert_eq!(channel_idx, 0);
4459 self.restart = true;
4460 }
4461 }
4462 }
4463
4464 async fn stop_workers(&mut self) {
4465 for worker in &mut self.workers {
4466 worker.stop().await;
4467 }
4468 }
4469
4470 async fn stop_primary_worker(&mut self) {
4471 self.workers[0].stop().await;
4472 }
4473
4474 async fn restore_guest_vf_state(&mut self, c_state: &mut CoordinatorState) {
4475 let primary = match self.primary_mut() {
4476 Some(primary) => primary,
4477 None => return,
4478 };
4479
4480 let virtual_function = c_state.virtual_function.as_mut();
4482 let guest_vf_id = match &virtual_function {
4483 Some(vf) => vf.id().await,
4484 None => None,
4485 };
4486 if let Some(guest_vf_id) = guest_vf_id {
4487 match primary.guest_vf_state {
4489 PrimaryChannelGuestVfState::AvailableAdvertised
4490 | PrimaryChannelGuestVfState::Restoring(
4491 saved_state::GuestVfState::AvailableAdvertised,
4492 ) => {
4493 if !primary.is_data_path_switched.unwrap_or(false) {
4494 let timer = PolledTimer::new(&c_state.adapter.driver);
4495 c_state.pending_vf_state = CoordinatorStatePendingVfState::Delay {
4496 timer,
4497 delay_until: Instant::now() + VF_DEVICE_DELAY,
4498 };
4499 }
4500 }
4501 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending { .. }
4502 | PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
4503 | PrimaryChannelGuestVfState::Ready
4504 | PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::Ready)
4505 | PrimaryChannelGuestVfState::DataPathSwitchPending { .. }
4506 | PrimaryChannelGuestVfState::Restoring(
4507 saved_state::GuestVfState::DataPathSwitchPending { .. },
4508 )
4509 | PrimaryChannelGuestVfState::DataPathSwitched
4510 | PrimaryChannelGuestVfState::Restoring(
4511 saved_state::GuestVfState::DataPathSwitched,
4512 )
4513 | PrimaryChannelGuestVfState::DataPathSynthetic => {
4514 c_state.pending_vf_state = CoordinatorStatePendingVfState::Pending;
4515 }
4516 _ => (),
4517 };
4518 if let PrimaryChannelGuestVfState::Restoring(
4520 saved_state::GuestVfState::DataPathSwitchPending {
4521 to_guest,
4522 id,
4523 result,
4524 },
4525 ) = primary.guest_vf_state
4526 {
4527 if result.is_some() {
4529 primary.guest_vf_state = PrimaryChannelGuestVfState::DataPathSwitchPending {
4530 to_guest,
4531 id,
4532 result,
4533 };
4534 return;
4535 }
4536 }
4537 primary.guest_vf_state = match primary.guest_vf_state {
4538 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending { .. }
4539 | PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
4540 | PrimaryChannelGuestVfState::DataPathSwitchPending { .. }
4541 | PrimaryChannelGuestVfState::Restoring(
4542 saved_state::GuestVfState::DataPathSwitchPending { .. },
4543 )
4544 | PrimaryChannelGuestVfState::DataPathSynthetic => {
4545 let (to_guest, id) = match primary.guest_vf_state {
4546 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending {
4547 to_guest,
4548 id,
4549 }
4550 | PrimaryChannelGuestVfState::DataPathSwitchPending {
4551 to_guest, id, ..
4552 }
4553 | PrimaryChannelGuestVfState::Restoring(
4554 saved_state::GuestVfState::DataPathSwitchPending {
4555 to_guest, id, ..
4556 },
4557 ) => (to_guest, id),
4558 _ => (true, None),
4559 };
4560 if matches!(
4564 c_state.pending_vf_state,
4565 CoordinatorStatePendingVfState::Delay { .. }
4566 ) {
4567 c_state.pending_vf_state = CoordinatorStatePendingVfState::Pending;
4568 }
4569 let result = c_state.endpoint.set_data_path_to_guest_vf(to_guest).await;
4570 let result = if let Err(err) = result {
4571 tracing::error!(
4572 err = err.as_ref() as &dyn std::error::Error,
4573 to_guest,
4574 "Failed to switch guest VF data path"
4575 );
4576 false
4577 } else {
4578 primary.is_data_path_switched = Some(to_guest);
4579 true
4580 };
4581 match primary.guest_vf_state {
4582 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending {
4583 ..
4584 }
4585 | PrimaryChannelGuestVfState::DataPathSwitchPending { .. }
4586 | PrimaryChannelGuestVfState::Restoring(
4587 saved_state::GuestVfState::DataPathSwitchPending { .. },
4588 ) => PrimaryChannelGuestVfState::DataPathSwitchPending {
4589 to_guest,
4590 id,
4591 result: Some(result),
4592 },
4593 _ if result => PrimaryChannelGuestVfState::DataPathSwitched,
4594 _ => PrimaryChannelGuestVfState::DataPathSynthetic,
4595 }
4596 }
4597 PrimaryChannelGuestVfState::Initializing
4598 | PrimaryChannelGuestVfState::Unavailable
4599 | PrimaryChannelGuestVfState::UnavailableFromAvailable
4600 | PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::NoState) => {
4601 PrimaryChannelGuestVfState::Available { vfid: guest_vf_id }
4602 }
4603 PrimaryChannelGuestVfState::AvailableAdvertised
4604 | PrimaryChannelGuestVfState::Restoring(
4605 saved_state::GuestVfState::AvailableAdvertised,
4606 ) => {
4607 if !primary.is_data_path_switched.unwrap_or(false) {
4608 PrimaryChannelGuestVfState::AvailableAdvertised
4609 } else {
4610 PrimaryChannelGuestVfState::DataPathSwitched
4613 }
4614 }
4615 PrimaryChannelGuestVfState::DataPathSwitched
4616 | PrimaryChannelGuestVfState::Restoring(
4617 saved_state::GuestVfState::DataPathSwitched,
4618 ) => PrimaryChannelGuestVfState::DataPathSwitched,
4619 PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::Ready) => {
4620 PrimaryChannelGuestVfState::Ready
4621 }
4622 _ => primary.guest_vf_state,
4623 };
4624 } else {
4625 match primary.guest_vf_state {
4627 PrimaryChannelGuestVfState::DataPathSwitchPending { to_guest, .. }
4628 | PrimaryChannelGuestVfState::Restoring(
4629 saved_state::GuestVfState::DataPathSwitchPending { to_guest, .. },
4630 ) => {
4631 if !to_guest {
4632 if let Err(err) = c_state.endpoint.set_data_path_to_guest_vf(false).await {
4633 tracing::warn!(
4634 err = err.as_ref() as &dyn std::error::Error,
4635 "Failed setting data path back to synthetic after guest VF was removed."
4636 );
4637 }
4638 primary.is_data_path_switched = Some(false);
4639 }
4640 }
4641 PrimaryChannelGuestVfState::DataPathSwitched
4642 | PrimaryChannelGuestVfState::Restoring(
4643 saved_state::GuestVfState::DataPathSwitched,
4644 ) => {
4645 if let Err(err) = c_state.endpoint.set_data_path_to_guest_vf(false).await {
4646 tracing::warn!(
4647 err = err.as_ref() as &dyn std::error::Error,
4648 "Failed setting data path back to synthetic after guest VF was removed."
4649 );
4650 }
4651 primary.is_data_path_switched = Some(false);
4652 }
4653 _ => (),
4654 }
4655 if let PrimaryChannelGuestVfState::AvailableAdvertised = primary.guest_vf_state {
4656 c_state.pending_vf_state = CoordinatorStatePendingVfState::Ready;
4657 }
4658 primary.guest_vf_state = match primary.guest_vf_state {
4660 PrimaryChannelGuestVfState::Initializing
4661 | PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::NoState)
4662 | PrimaryChannelGuestVfState::Available { .. } => {
4663 PrimaryChannelGuestVfState::Unavailable
4664 }
4665 PrimaryChannelGuestVfState::AvailableAdvertised
4666 | PrimaryChannelGuestVfState::Restoring(
4667 saved_state::GuestVfState::AvailableAdvertised,
4668 )
4669 | PrimaryChannelGuestVfState::Ready
4670 | PrimaryChannelGuestVfState::Restoring(saved_state::GuestVfState::Ready) => {
4671 PrimaryChannelGuestVfState::UnavailableFromAvailable
4672 }
4673 PrimaryChannelGuestVfState::DataPathSwitchPending { to_guest, id, .. }
4674 | PrimaryChannelGuestVfState::Restoring(
4675 saved_state::GuestVfState::DataPathSwitchPending { to_guest, id, .. },
4676 ) => PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending {
4677 to_guest,
4678 id,
4679 },
4680 PrimaryChannelGuestVfState::DataPathSwitched
4681 | PrimaryChannelGuestVfState::Restoring(
4682 saved_state::GuestVfState::DataPathSwitched,
4683 )
4684 | PrimaryChannelGuestVfState::DataPathSynthetic => {
4685 PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
4686 }
4687 _ => primary.guest_vf_state,
4688 }
4689 }
4690 }
4691
4692 async fn restart_queues(&mut self, c_state: &mut CoordinatorState) -> Result<(), WorkerError> {
4693 let (num_queues, active_queues, active_queue_count) = if let Some(state) = self.workers[0]
4697 .state()
4698 .and_then(|worker| worker.state.ready())
4699 {
4700 let num_queues = state.state.primary.as_ref().unwrap().requested_num_queues;
4701 let mut active_queues = Vec::new();
4702 let active_queue_count = if let Some(rss_state) =
4703 state.state.primary.as_ref().unwrap().rss_state.as_ref()
4704 {
4705 active_queues.clone_from(&rss_state.indirection_table);
4706 active_queues.sort();
4707 active_queues.dedup();
4708 active_queues = active_queues
4709 .into_iter()
4710 .filter(|&index| index < num_queues)
4711 .collect::<Vec<_>>();
4712 if !active_queues.is_empty() {
4713 active_queues.len() as u16
4714 } else {
4715 tracelimit::warn_ratelimited!(
4716 num_queues,
4717 indirection_table_len = rss_state.indirection_table.len(),
4718 "RSS indirection table has no entries within the valid queue range, falling back to num_queues",
4719 );
4720 num_queues
4721 }
4722 } else {
4723 num_queues
4724 };
4725
4726 RxBufferRanges::validate_params(
4727 state.buffers.recv_buffer.count,
4728 active_queue_count.into(),
4729 )?;
4730
4731 GuestBuffers::validate_config(
4732 &state.buffers.recv_buffer.gpadl,
4733 state.buffers.recv_buffer.sub_allocation_size,
4734 state.buffers.ndis_config.mtu,
4735 )
4736 .map_err(WorkerError::GuestBuffers)?;
4737
4738 (num_queues, active_queues, active_queue_count)
4739 } else {
4740 (0, Vec::new(), 0)
4742 };
4743
4744 let drivers = self
4746 .workers
4747 .iter_mut()
4748 .map(|worker| {
4749 let task = worker.task_mut();
4750 task.queue_state = None;
4751 task.driver.clone()
4752 })
4753 .collect::<Vec<_>>();
4754
4755 c_state.endpoint.stop().await;
4756
4757 let (primary_worker, subworkers) = if let [primary, sub @ ..] = self.workers.as_mut_slice()
4758 {
4759 (primary, sub)
4760 } else {
4761 unreachable!()
4762 };
4763
4764 let state = primary_worker
4765 .state_mut()
4766 .and_then(|worker| worker.state.ready_mut());
4767
4768 let state = if let Some(state) = state {
4769 state
4770 } else {
4771 return Ok(());
4772 };
4773
4774 self.buffers = Some(state.buffers.clone());
4776
4777 let (ranges, mut remote_buffer_id_recvs) =
4779 RxBufferRanges::new(state.buffers.recv_buffer.count, active_queue_count.into())?;
4780 let ranges = Arc::new(ranges);
4781
4782 let mut queues = Vec::new();
4783 let mut rx_buffers = Vec::new();
4784 let mut per_queue_rx: Vec<Vec<RxId>> = Vec::new();
4785 let guest_buffers;
4786 {
4787 let buffers = &state.buffers;
4788 guest_buffers = Arc::new(
4789 GuestBuffers::new(
4790 buffers.mem.clone(),
4791 buffers.recv_buffer.gpadl.clone(),
4792 buffers.recv_buffer.sub_allocation_size,
4793 buffers.ndis_config.mtu,
4794 )
4795 .map_err(WorkerError::GuestBuffers)?,
4796 );
4797
4798 let mut queue_config = Vec::new();
4801 let initial_rx;
4802 {
4803 let states = std::iter::once(Some(&*state)).chain(
4804 subworkers
4805 .iter()
4806 .map(|worker| worker.state().and_then(|worker| worker.state.ready())),
4807 );
4808
4809 initial_rx = (RX_RESERVED_CONTROL_BUFFERS..state.buffers.recv_buffer.count)
4810 .filter(|&n| states.clone().flatten().all(|s| s.state.rx_bufs.is_free(n)))
4811 .map(RxId)
4812 .collect::<Vec<_>>();
4813
4814 let mut initial_rx = initial_rx.as_slice();
4815 let mut range_start = 0;
4816 let primary_queue_excluded = !active_queues.is_empty() && active_queues[0] != 0;
4817 let first_queue = if !primary_queue_excluded {
4818 0
4819 } else {
4820 queue_config.push(QueueConfig {
4824 driver: Box::new(drivers[0].clone()),
4825 });
4826 per_queue_rx.push(Vec::new());
4827 rx_buffers.push(RxBufferRange::new(
4828 ranges.clone(),
4829 0..RX_RESERVED_CONTROL_BUFFERS,
4830 None,
4831 ));
4832 range_start = RX_RESERVED_CONTROL_BUFFERS;
4833 1
4834 };
4835 for queue_index in first_queue..num_queues {
4836 let queue_active = active_queues.is_empty()
4837 || active_queues.binary_search(&queue_index).is_ok();
4838 let (range_end, end, buffer_id_recv) = if queue_active {
4839 let range_end = if rx_buffers.len() as u16 == active_queue_count - 1 {
4840 state.buffers.recv_buffer.count
4842 } else if queue_index == 0 {
4843 RX_RESERVED_CONTROL_BUFFERS + ranges.buffers_per_queue
4845 } else {
4846 range_start + ranges.buffers_per_queue
4847 };
4848 (
4849 range_end,
4850 initial_rx.partition_point(|id| id.0 < range_end),
4851 Some(remote_buffer_id_recvs.remove(0)),
4852 )
4853 } else {
4854 (range_start, 0, None)
4855 };
4856
4857 let (this, rest) = initial_rx.split_at(end);
4858 queue_config.push(QueueConfig {
4859 driver: Box::new(drivers[queue_index as usize].clone()),
4860 });
4861 per_queue_rx.push(this.to_vec());
4862 initial_rx = rest;
4863 rx_buffers.push(RxBufferRange::new(
4864 ranges.clone(),
4865 range_start..range_end,
4866 buffer_id_recv,
4867 ));
4868
4869 range_start = range_end;
4870 }
4871 }
4872
4873 let primary = state.state.primary.as_mut().unwrap();
4874 tracing::debug!(num_queues, "enabling endpoint");
4875
4876 let rss = primary
4877 .rss_state
4878 .as_ref()
4879 .map(|rss| net_backend::RssConfig {
4880 key: &rss.key,
4881 indirection_table: &rss.indirection_table,
4882 flags: 0,
4883 });
4884
4885 c_state
4886 .endpoint
4887 .get_queues(queue_config, rss.as_ref(), &mut queues)
4888 .instrument(tracing::info_span!("netvsp_get_queues"))
4889 .await
4890 .map_err(WorkerError::Endpoint)?;
4891
4892 assert_eq!(queues.len(), num_queues as usize);
4893
4894 self.channel_control
4896 .enable_subchannels(num_queues - 1)
4897 .expect("already validated");
4898
4899 self.num_queues = num_queues;
4900 }
4901
4902 let packet_size = state.buffers.version.into();
4904 self.active_packet_filter = self.workers[0].state().unwrap().channel.packet_filter;
4905 for (((worker, mut queue), rx_buffer), initial) in self
4907 .workers
4908 .iter_mut()
4909 .zip(queues)
4910 .zip(rx_buffers)
4911 .zip(per_queue_rx)
4912 {
4913 let mut pool = BufferPool::new(guest_buffers.clone());
4914 if !initial.is_empty() {
4915 queue.rx_avail(&mut pool, &initial);
4916 }
4917 worker.task_mut().queue_state = Some(QueueState {
4918 queue,
4919 pool,
4920 target_vp_set: false,
4921 rx_buffer_range: rx_buffer,
4922 });
4923 if let Some(worker) = worker.state_mut() {
4925 worker.channel.packet_filter = self.active_packet_filter;
4926 if let Some(ready_state) = worker.state.ready_mut() {
4929 ready_state.state.pending_rx_packets.clear();
4930 ready_state.reset_tx_after_endpoint_stop();
4931 }
4932
4933 worker.channel.packet_size = packet_size;
4936 }
4937 }
4938
4939 Ok(())
4940 }
4941
4942 fn primary_mut(&mut self) -> Option<&mut PrimaryChannelState> {
4943 self.workers[0]
4944 .state_mut()
4945 .unwrap()
4946 .state
4947 .ready_mut()?
4948 .state
4949 .primary
4950 .as_mut()
4951 }
4952
4953 async fn update_guest_vf_state(&mut self, c_state: &mut CoordinatorState) {
4954 self.stop_primary_worker().await;
4955 self.restore_guest_vf_state(c_state).await;
4956 }
4957}
4958
4959impl<T: RingMem + 'static + Sync> AsyncRun<Worker<T>> for NetQueue {
4960 async fn run(
4961 &mut self,
4962 stop: &mut StopTask<'_>,
4963 worker: &mut Worker<T>,
4964 ) -> Result<(), task_control::Cancelled> {
4965 match worker.process(stop, self).await {
4966 Ok(()) | Err(WorkerError::BufferRevoked) => {}
4967 Err(WorkerError::Cancelled(cancelled)) => return Err(cancelled),
4968 Err(err) => {
4969 tracing::error!(
4970 error = &err as &dyn std::error::Error,
4971 channel_idx = worker.channel_idx,
4972 "netvsp error"
4973 );
4974 }
4975 }
4976 Ok(())
4977 }
4978}
4979
4980impl<T: RingMem + 'static> Worker<T> {
4981 async fn process(
4982 &mut self,
4983 stop: &mut StopTask<'_>,
4984 queue: &mut NetQueue,
4985 ) -> Result<(), WorkerError> {
4986 loop {
4990 match &mut self.state {
4991 WorkerState::Init(initializing) => {
4992 assert_eq!(self.channel_idx, 0);
4993
4994 tracelimit::info_ratelimited!("network accepted");
4995
4996 let (buffers, state) = stop
4997 .until_stopped(self.channel.initialize(initializing, self.mem.clone()))
4998 .await??;
4999
5000 let state = ReadyState {
5001 buffers: Arc::new(buffers),
5002 state,
5003 data: ProcessingData::new(),
5004 };
5005
5006 if let Err(err) = self
5008 .coordinator_send
5009 .try_send(CoordinatorMessage::Restart { channel_idx: 0 })
5010 {
5011 tracelimit::error_ratelimited!(
5012 error = &err as &dyn std::error::Error,
5013 channel_idx = self.channel_idx,
5014 "failed to send restart message to coordinator"
5015 );
5016 }
5017
5018 tracelimit::info_ratelimited!("network initialized");
5019 self.state = WorkerState::WaitingForCoordinator(Some(state));
5020 }
5021 WorkerState::WaitingForCoordinator(_) => {
5022 assert_eq!(self.channel_idx, 0);
5023 stop.until_stopped(pending()).await?
5026 }
5027 WorkerState::Ready(state) => {
5028 let queue_state = if let Some(queue_state) = &mut queue.queue_state {
5029 if !queue_state.target_vp_set {
5030 if let Some(target_vp) = self.target_vp {
5031 tracing::debug!(
5032 channel_idx = self.channel_idx,
5033 target_vp,
5034 "updating target VP"
5035 );
5036 queue_state.queue.update_target_vp(target_vp).await;
5037 queue_state.target_vp_set = true;
5038 }
5039 }
5040
5041 queue_state
5042 } else {
5043 stop.until_stopped(pending()).await?
5045 };
5046
5047 let result = self.channel.main_loop(stop, state, queue_state).await;
5048 let msg = match result {
5049 Ok(restart) => {
5050 assert_eq!(self.channel_idx, 0);
5051 restart
5052 }
5053 Err(WorkerError::EndpointRequiresQueueRestart(err)) => {
5054 tracelimit::warn_ratelimited!(
5055 err = err.as_ref() as &dyn std::error::Error,
5056 channel_idx = self.channel_idx,
5057 "Endpoint requires queues to restart",
5058 );
5059 CoordinatorMessage::Restart {
5060 channel_idx: self.channel_idx,
5061 }
5062 }
5063 Err(err) => return Err(err),
5064 };
5065
5066 if self.channel_idx == 0 {
5069 let WorkerState::Ready(ready) = std::mem::replace(
5070 &mut self.state,
5071 WorkerState::WaitingForCoordinator(None),
5072 ) else {
5073 unreachable!("must be running in ready state")
5074 };
5075 let _ = std::mem::replace(
5076 &mut self.state,
5077 WorkerState::WaitingForCoordinator(Some(ready)),
5078 );
5079 }
5080 self.coordinator_send
5081 .try_send(msg)
5082 .map_err(WorkerError::CoordinatorMessageSendFailed)?;
5083 stop.until_stopped(pending()).await?
5084 }
5085 }
5086 }
5087 }
5088}
5089
5090impl<T: 'static + RingMem> NetChannel<T> {
5091 fn try_next_packet<'a>(
5092 &mut self,
5093 send_buffer: Option<&SendBuffer>,
5094 version: Option<Version>,
5095 external_data: &'a mut MultiPagedRangeBuf,
5096 ) -> Result<Option<Packet<'a>>, WorkerError> {
5097 let (mut read, _) = self.queue.split();
5098 let packet = match read.try_read() {
5099 Ok(packet) => parse_packet(&packet, send_buffer, version, external_data)
5100 .map_err(WorkerError::Packet)?,
5101 Err(queue::TryReadError::Empty) => return Ok(None),
5102 Err(queue::TryReadError::Queue(err)) => return Err(err.into()),
5103 };
5104
5105 tracing::trace!(target: "netvsp/vmbus", data = ?packet.data, "incoming vmbus packet");
5106 Ok(Some(packet))
5107 }
5108
5109 async fn next_packet<'a>(
5110 &mut self,
5111 send_buffer: Option<&'a SendBuffer>,
5112 version: Option<Version>,
5113 external_data: &'a mut MultiPagedRangeBuf,
5114 ) -> Result<Packet<'a>, WorkerError> {
5115 let (mut read, _) = self.queue.split();
5116 let mut packet_ref = read.read().await?;
5117 let packet = parse_packet(&packet_ref, send_buffer, version, external_data)
5118 .map_err(WorkerError::Packet)?;
5119 if matches!(packet.data, PacketData::RndisPacket(_)) {
5120 tracing::trace!(target: "netvsp/vmbus", "detected rndis initialization message");
5122 packet_ref.revert();
5123 }
5124 tracing::trace!(target: "netvsp/vmbus", data = ?packet.data, "incoming vmbus packet");
5125 Ok(packet)
5126 }
5127
5128 fn is_ready_to_initialize(initializing: &InitState, allow_missing_send_buffer: bool) -> bool {
5129 (initializing.ndis_config.is_some() || initializing.version < Version::V2)
5130 && initializing.ndis_version.is_some()
5131 && (initializing.send_buffer.is_some() || allow_missing_send_buffer)
5132 && initializing.recv_buffer.is_some()
5133 }
5134
5135 async fn initialize(
5136 &mut self,
5137 initializing: &mut Option<InitState>,
5138 mem: GuestMemory,
5139 ) -> Result<(ChannelBuffers, ActiveState), WorkerError> {
5140 let mut has_init_packet_arrived = false;
5141 loop {
5142 if let Some(initializing) = &mut *initializing {
5143 if Self::is_ready_to_initialize(initializing, false) || has_init_packet_arrived {
5144 let recv_buffer = initializing.recv_buffer.take().unwrap();
5145 let send_buffer = initializing.send_buffer.take();
5146 let state = ActiveState::new(
5147 Some(PrimaryChannelState::new(
5148 self.adapter.offload_support.clone(),
5149 )),
5150 recv_buffer.count,
5151 );
5152 let buffers = ChannelBuffers {
5153 version: initializing.version,
5154 mem,
5155 recv_buffer,
5156 send_buffer,
5157 ndis_version: initializing.ndis_version.take().unwrap(),
5158 ndis_config: initializing.ndis_config.take().unwrap_or(NdisConfig {
5159 mtu: DEFAULT_MTU,
5160 capabilities: protocol::NdisConfigCapabilities::new(),
5161 }),
5162 };
5163
5164 break Ok((buffers, state));
5165 }
5166 }
5167
5168 self.queue
5171 .split()
5172 .1
5173 .wait_ready(ring::PacketSize::completion(protocol::PACKET_SIZE_V61))
5174 .await?;
5175
5176 let mut external_data = MultiPagedRangeBuf::new();
5177 let packet = self
5178 .next_packet(
5179 None,
5180 initializing.as_ref().map(|x| x.version),
5181 &mut external_data,
5182 )
5183 .await?;
5184
5185 if let Some(initializing) = &mut *initializing {
5186 match packet.data {
5187 PacketData::SendNdisConfig(config) => {
5188 if initializing.ndis_config.is_some() {
5189 return Err(WorkerError::UnexpectedPacketOrder(
5190 PacketOrderError::SendNdisConfigExists,
5191 ));
5192 }
5193
5194 let mtu = if config.mtu >= MIN_MTU && config.mtu <= MAX_MTU {
5196 config.mtu
5197 } else {
5198 DEFAULT_MTU
5199 };
5200
5201 self.send_completion(packet.transaction_id, None)?;
5203 initializing.ndis_config = Some(NdisConfig {
5204 mtu,
5205 capabilities: config.capabilities,
5206 });
5207 }
5208 PacketData::SendNdisVersion(version) => {
5209 if initializing.ndis_version.is_some() {
5210 return Err(WorkerError::UnexpectedPacketOrder(
5211 PacketOrderError::SendNdisVersionExists,
5212 ));
5213 }
5214
5215 self.send_completion(packet.transaction_id, None)?;
5217 initializing.ndis_version = Some(NdisVersion {
5218 major: version.ndis_major_version,
5219 minor: version.ndis_minor_version,
5220 });
5221 }
5222 PacketData::SendReceiveBuffer(message) => {
5223 if initializing.recv_buffer.is_some() {
5224 return Err(WorkerError::UnexpectedPacketOrder(
5225 PacketOrderError::SendReceiveBufferExists,
5226 ));
5227 }
5228
5229 let mtu = if let Some(cfg) = &initializing.ndis_config {
5230 cfg.mtu
5231 } else if initializing.version < Version::V2 {
5232 DEFAULT_MTU
5233 } else {
5234 return Err(WorkerError::UnexpectedPacketOrder(
5235 PacketOrderError::SendReceiveBufferMissingMTU,
5236 ));
5237 };
5238
5239 let sub_allocation_size = sub_allocation_size_for_mtu(mtu);
5240
5241 let recv_buffer = ReceiveBuffer::new(
5242 &self.gpadl_map,
5243 message.gpadl_handle,
5244 message.id,
5245 sub_allocation_size,
5246 )?;
5247
5248 self.send_completion(
5249 packet.transaction_id,
5250 Some(&self.message(
5251 protocol::MESSAGE1_TYPE_SEND_RECEIVE_BUFFER_COMPLETE,
5252 protocol::Message1SendReceiveBufferComplete {
5253 status: protocol::Status::SUCCESS,
5254 num_sections: 1,
5255 sections: [protocol::ReceiveBufferSection {
5256 offset: 0,
5257 sub_allocation_size: recv_buffer.sub_allocation_size,
5258 num_sub_allocations: recv_buffer.count,
5259 end_offset: recv_buffer.sub_allocation_size
5260 * recv_buffer.count,
5261 }],
5262 },
5263 )),
5264 )?;
5265 initializing.recv_buffer = Some(recv_buffer);
5266 }
5267 PacketData::SendSendBuffer(message) => {
5268 if initializing.send_buffer.is_some() {
5269 return Err(WorkerError::UnexpectedPacketOrder(
5270 PacketOrderError::SendSendBufferExists,
5271 ));
5272 }
5273
5274 let send_buffer = SendBuffer::new(&self.gpadl_map, message.gpadl_handle)?;
5275 self.send_completion(
5276 packet.transaction_id,
5277 Some(&self.message(
5278 protocol::MESSAGE1_TYPE_SEND_SEND_BUFFER_COMPLETE,
5279 protocol::Message1SendSendBufferComplete {
5280 status: protocol::Status::SUCCESS,
5281 section_size: 6144,
5282 },
5283 )),
5284 )?;
5285
5286 initializing.send_buffer = Some(send_buffer);
5287 }
5288 PacketData::RndisPacket(rndis_packet) => {
5289 if !Self::is_ready_to_initialize(initializing, true) {
5290 return Err(WorkerError::UnexpectedPacketOrder(
5291 PacketOrderError::UnexpectedRndisPacket,
5292 ));
5293 }
5294 tracing::debug!(
5295 channel_type = rndis_packet.channel_type,
5296 "RndisPacket received during initialization, assuming MESSAGE_TYPE_INITIALIZE_MSG"
5297 );
5298 has_init_packet_arrived = true;
5299 }
5300 _ => {
5301 return Err(WorkerError::UnexpectedPacketOrder(
5302 PacketOrderError::InvalidPacketData,
5303 ));
5304 }
5305 }
5306 } else {
5307 match packet.data {
5308 PacketData::Init(init) => {
5309 let requested_version = init.protocol_version;
5310 let version = check_version(requested_version);
5311 let mut data = protocol::MessageInitComplete {
5312 deprecated: protocol::INVALID_PROTOCOL_VERSION,
5313 maximum_mdl_chain_length: 34,
5314 status: protocol::Status::NONE,
5315 };
5316 if let Some(version) = version {
5317 if version == Version::V1 {
5318 data.deprecated = Version::V1 as u32;
5319 }
5320 data.status = protocol::Status::SUCCESS;
5321 } else {
5322 tracing::debug!(requested_version, "unrecognized version");
5323 }
5324 let message = self.message(protocol::MESSAGE_TYPE_INIT_COMPLETE, data);
5325 self.send_completion(packet.transaction_id, Some(&message))?;
5326
5327 if let Some(version) = version {
5328 tracelimit::info_ratelimited!(?version, "network negotiated");
5329
5330 self.packet_size = version.into();
5332
5333 *initializing = Some(InitState {
5334 version,
5335 ndis_config: None,
5336 ndis_version: None,
5337 recv_buffer: None,
5338 send_buffer: None,
5339 });
5340 }
5341 }
5342 _ => unreachable!(),
5343 }
5344 }
5345 }
5346 }
5347
5348 async fn main_loop(
5349 &mut self,
5350 stop: &mut StopTask<'_>,
5351 ready_state: &mut ReadyState,
5352 queue_state: &mut QueueState,
5353 ) -> Result<CoordinatorMessage, WorkerError> {
5354 let buffers = &ready_state.buffers;
5355 let state = &mut ready_state.state;
5356 let data = &mut ready_state.data;
5357
5358 let ring_spare_capacity = {
5359 let (_, send) = self.queue.split();
5360 let mut limit = if self.can_use_ring_size_opt {
5361 self.adapter.ring_size_limit.load(Ordering::Relaxed)
5362 } else {
5363 0
5364 };
5365 if limit == 0 {
5366 limit = send.capacity() - 2048;
5367 }
5368 send.capacity() - limit
5369 };
5370
5371 if !state.pending_rx_packets.is_empty()
5373 && self.packet_filter != rndisprot::NDIS_PACKET_TYPE_NONE
5374 {
5375 let (front, back) = state.pending_rx_packets.as_slices();
5376 queue_state.queue.rx_avail(&mut queue_state.pool, front);
5377 queue_state.queue.rx_avail(&mut queue_state.pool, back);
5378 state.pending_rx_packets.clear();
5379 }
5380
5381 if let Some(primary) = state.primary.as_mut() {
5383 if primary.requested_num_queues > 1 && !primary.tx_spread_sent {
5384 let num_channels_opened =
5385 self.adapter.num_sub_channels_opened.load(Ordering::Relaxed);
5386 if num_channels_opened == primary.requested_num_queues as usize {
5387 let (_, mut send) = self.queue.split();
5388 stop.until_stopped(send.wait_ready(MIN_STATE_CHANGE_RING_SIZE))
5389 .await??;
5390 self.guest_send_indirection_table(buffers.version, num_channels_opened as u32);
5391 primary.tx_spread_sent = true;
5392 }
5393 }
5394 if let PendingLinkAction::Active(up) = primary.pending_link_action {
5395 let (_, mut send) = self.queue.split();
5396 stop.until_stopped(send.wait_ready(MIN_STATE_CHANGE_RING_SIZE))
5397 .await??;
5398 if let Some(id) = primary.free_control_buffers.pop() {
5399 let connect = if primary.guest_link_up != up {
5400 primary.pending_link_action = PendingLinkAction::Default;
5401 up
5402 } else {
5403 primary.pending_link_action =
5406 PendingLinkAction::Delay(primary.guest_link_up);
5407 !primary.guest_link_up
5408 };
5409 assert!(state.rx_bufs.is_free(id.0));
5411 state.rx_bufs.allocate(std::iter::once(id.0)).unwrap();
5412 let state_to_send = if connect {
5413 rndisprot::STATUS_MEDIA_CONNECT
5414 } else {
5415 rndisprot::STATUS_MEDIA_DISCONNECT
5416 };
5417 tracing::info!(
5418 connect,
5419 mac_address = %self.adapter.mac_address,
5420 "sending link status"
5421 );
5422
5423 self.indicate_status(buffers, id.0, state_to_send, &[])?;
5424 primary.guest_link_up = connect;
5425 } else {
5426 primary.pending_link_action = PendingLinkAction::Delay(up);
5427 }
5428
5429 match primary.pending_link_action {
5430 PendingLinkAction::Delay(_) => {
5431 return Ok(CoordinatorMessage::StartTimer(
5432 Instant::now() + LINK_DELAY_DURATION,
5433 ));
5434 }
5435 PendingLinkAction::Active(_) => panic!("State should not be Active"),
5436 _ => {}
5437 }
5438 }
5439 match primary.guest_vf_state {
5440 PrimaryChannelGuestVfState::Available { .. }
5441 | PrimaryChannelGuestVfState::UnavailableFromAvailable
5442 | PrimaryChannelGuestVfState::UnavailableFromDataPathSwitchPending { .. }
5443 | PrimaryChannelGuestVfState::UnavailableFromDataPathSwitched
5444 | PrimaryChannelGuestVfState::DataPathSwitchPending { .. }
5445 | PrimaryChannelGuestVfState::DataPathSynthetic => {
5446 let (_, mut send) = self.queue.split();
5447 stop.until_stopped(send.wait_ready(MIN_STATE_CHANGE_RING_SIZE))
5448 .await??;
5449 if let Some(message) = self.handle_state_change(primary, buffers).await? {
5450 return Ok(message);
5451 }
5452 }
5453 _ => (),
5454 }
5455 }
5456
5457 loop {
5458 let ring_full = {
5463 let (_, mut send) = self.queue.split();
5464 !send.can_write(ring_spare_capacity)?
5465 };
5466
5467 let did_some_work = (!ring_full
5468 && self.process_endpoint_rx(
5469 buffers,
5470 state,
5471 data,
5472 queue_state.queue.as_mut(),
5473 &mut queue_state.pool,
5474 )?)
5475 | self.process_ring_buffer(buffers, state, data, queue_state)?
5476 | (!ring_full
5477 && self.process_endpoint_tx(
5478 state,
5479 data,
5480 queue_state.queue.as_mut(),
5481 &mut queue_state.pool,
5482 )?)
5483 | self.transmit_pending_segments(state, data, queue_state)?
5484 | self.send_pending_packets(state)?;
5485
5486 if !did_some_work {
5487 state.stats.spurious_wakes.increment();
5488 }
5489
5490 self.process_control_messages(buffers, state)?;
5493
5494 let restart = stop
5498 .until_stopped(std::future::poll_fn(
5499 |cx| -> Poll<Option<CoordinatorMessage>> {
5500 if !ring_full {
5504 if queue_state
5506 .queue
5507 .poll_ready(cx, &mut queue_state.pool)
5508 .is_ready()
5509 {
5510 tracing::trace!("endpoint ready");
5511 return Poll::Ready(None);
5512 }
5513 }
5514
5515 let (mut recv, mut send) = self.queue.split();
5518 if state.free_tx_packets.len() >= self.adapter.free_tx_packet_threshold
5519 && data.tx_segments.is_empty()
5520 && recv.poll_ready(cx).is_ready()
5521 {
5522 tracing::trace!("incoming ring ready");
5523 return Poll::Ready(None);
5524 }
5525
5526 let mut pending_send_size = self.pending_send_size;
5533 if ring_full {
5534 pending_send_size = ring_spare_capacity;
5535 }
5536 if pending_send_size != 0
5537 && send.poll_ready(cx, pending_send_size).is_ready()
5538 {
5539 tracing::trace!("outgoing ring ready");
5540 return Poll::Ready(None);
5541 }
5542
5543 if let Some(remote_buffer_id_recv) =
5549 &mut queue_state.rx_buffer_range.remote_buffer_id_recv
5550 {
5551 while let Poll::Ready(Some(id)) =
5552 remote_buffer_id_recv.poll_next_unpin(cx)
5553 {
5554 if id >= RX_RESERVED_CONTROL_BUFFERS {
5555 queue_state
5556 .queue
5557 .rx_avail(&mut queue_state.pool, &[RxId(id)]);
5558 } else {
5559 state
5560 .primary
5561 .as_mut()
5562 .unwrap()
5563 .free_control_buffers
5564 .push(ControlMessageId(id));
5565 }
5566 }
5567 }
5568
5569 if let Some(restart) = self.restart.take() {
5570 return Poll::Ready(Some(restart));
5571 }
5572
5573 tracing::trace!("network waiting");
5574 Poll::Pending
5575 },
5576 ))
5577 .await?;
5578
5579 if let Some(restart) = restart {
5580 break Ok(restart);
5581 }
5582 }
5583 }
5584
5585 fn process_endpoint_rx(
5586 &mut self,
5587 buffers: &ChannelBuffers,
5588 state: &mut ActiveState,
5589 data: &mut ProcessingData,
5590 epqueue: &mut dyn net_backend::Queue,
5591 pool: &mut BufferPool,
5592 ) -> Result<bool, WorkerError> {
5593 let n = epqueue
5594 .rx_poll(pool, &mut data.rx_ready)
5595 .map_err(WorkerError::Endpoint)?;
5596
5597 if n == 0 {
5598 return Ok(false);
5599 }
5600
5601 state.stats.rx_packets_per_wake.add_sample(n as u64);
5602 state.stats.rx_vlan_packets.add(pool.take_rx_vlan_count());
5603
5604 if self.packet_filter == rndisprot::NDIS_PACKET_TYPE_NONE {
5605 tracing::trace!(
5606 packet_filter = self.packet_filter,
5607 "rx packets dropped due to packet filter"
5608 );
5609 state.pending_rx_packets.extend(&data.rx_ready[..n]);
5614 state.stats.rx_dropped_filtered.add(n as u64);
5615 return Ok(false);
5616 }
5617
5618 let transaction_id = data.rx_ready[0].0.into();
5619 let ready_ids = data.rx_ready[..n].iter().map(|&RxId(id)| id);
5620
5621 state.rx_bufs.allocate(ready_ids.clone()).unwrap();
5622
5623 let len = buffers.recv_buffer.sub_allocation_size as usize;
5626 data.transfer_pages.clear();
5627 data.transfer_pages
5628 .extend(ready_ids.map(|id| buffers.recv_buffer.transfer_page_range(id, len)));
5629
5630 match self.try_send_rndis_message(
5631 transaction_id,
5632 protocol::DATA_CHANNEL_TYPE,
5633 buffers.recv_buffer.id,
5634 &data.transfer_pages,
5635 )? {
5636 None => {
5637 state.stats.rx_packets.add(n as u64);
5639 }
5640 Some(_) => {
5641 state.stats.rx_dropped_ring_full.add(n as u64);
5645
5646 state.rx_bufs.free(data.rx_ready[0].0);
5647 epqueue.rx_avail(pool, &data.rx_ready[..n]);
5648 }
5649 }
5650
5651 Ok(true)
5652 }
5653
5654 fn process_endpoint_tx(
5655 &mut self,
5656 state: &mut ActiveState,
5657 data: &mut ProcessingData,
5658 epqueue: &mut dyn net_backend::Queue,
5659 pool: &mut BufferPool,
5660 ) -> Result<bool, WorkerError> {
5661 let result = epqueue.tx_poll(pool, &mut data.tx_done);
5663
5664 match result {
5665 Ok(n) => {
5666 if n == 0 {
5667 return Ok(false);
5668 }
5669
5670 for &id in &data.tx_done[..n] {
5671 let tx_packet = &mut state.pending_tx_packets[id.0 as usize];
5672 assert!(tx_packet.pending_packet_count > 0);
5673 tx_packet.pending_packet_count -= 1;
5674 if tx_packet.pending_packet_count == 0 {
5675 self.complete_tx_packet(state, id, protocol::Status::SUCCESS)?;
5676 }
5677 }
5678
5679 Ok(true)
5680 }
5681 Err(TxError::TryRestart(err)) => {
5682 Err(WorkerError::EndpointRequiresQueueRestart(err))
5685 }
5686 Err(TxError::Fatal(err)) => Err(WorkerError::Endpoint(err)),
5687 }
5688 }
5689
5690 fn switch_data_path(
5691 &mut self,
5692 state: &mut ActiveState,
5693 use_guest_vf: bool,
5694 transaction_id: Option<u64>,
5695 ) -> Result<(), WorkerError> {
5696 let primary = state.primary.as_mut().unwrap();
5697 let mut queue_switch_operation = false;
5698 match primary.guest_vf_state {
5699 PrimaryChannelGuestVfState::AvailableAdvertised | PrimaryChannelGuestVfState::Ready => {
5700 if use_guest_vf || primary.is_data_path_switched.is_none() {
5709 primary.guest_vf_state = PrimaryChannelGuestVfState::DataPathSwitchPending {
5710 to_guest: use_guest_vf,
5711 id: transaction_id,
5712 result: None,
5713 };
5714 queue_switch_operation = true;
5715 }
5716 }
5717 PrimaryChannelGuestVfState::DataPathSwitched => {
5718 if !use_guest_vf {
5719 primary.guest_vf_state = PrimaryChannelGuestVfState::DataPathSwitchPending {
5720 to_guest: false,
5721 id: transaction_id,
5722 result: None,
5723 };
5724 queue_switch_operation = true;
5725 }
5726 }
5727 _ if use_guest_vf => {
5728 tracing::warn!(
5729 state = %primary.guest_vf_state,
5730 use_guest_vf,
5731 "Data path switch requested while device is in wrong state"
5732 );
5733 }
5734 _ => (),
5735 };
5736 if queue_switch_operation {
5737 self.send_coordinator_update_vf();
5738 } else {
5739 self.send_completion(transaction_id, None)?;
5740 }
5741 Ok(())
5742 }
5743
5744 fn process_ring_buffer(
5745 &mut self,
5746 buffers: &ChannelBuffers,
5747 state: &mut ActiveState,
5748 data: &mut ProcessingData,
5749 queue_state: &mut QueueState,
5750 ) -> Result<bool, WorkerError> {
5751 if !data.tx_segments.is_empty() {
5752 return Ok(false);
5756 }
5757 let mut total_packets = 0;
5758 let mut did_some_work = false;
5759 loop {
5760 if state.free_tx_packets.is_empty() {
5761 break;
5762 }
5763 let packet = if let Some(packet) = self.try_next_packet(
5764 buffers.send_buffer.as_ref(),
5765 Some(buffers.version),
5766 &mut data.external_data,
5767 )? {
5768 packet
5769 } else {
5770 break;
5771 };
5772
5773 did_some_work = true;
5774 match packet.data {
5775 PacketData::RndisPacket(_) => {
5776 let id = state.free_tx_packets.pop().unwrap();
5777 let result: Result<usize, WorkerError> =
5778 self.handle_rndis(buffers, id, state, &packet, &mut data.tx_segments);
5779 match result {
5780 Ok(num_packets) => {
5781 total_packets += num_packets as u64;
5782 if num_packets == 0 {
5783 self.complete_tx_packet(state, id, protocol::Status::SUCCESS)?;
5784 }
5785 }
5786 Err(err) => {
5787 tracelimit::error_ratelimited!(
5788 error = &err as &dyn std::error::Error,
5789 "failed to handle RNDIS packet"
5790 );
5791 self.complete_tx_packet(state, id, protocol::Status::FAILURE)?;
5792 }
5793 };
5794 }
5795 PacketData::RndisPacketComplete(_completion) => {
5796 data.rx_done.clear();
5797 state
5798 .release_recv_buffers(
5799 packet
5800 .transaction_id
5801 .expect("completion packets have transaction id by construction"),
5802 &queue_state.rx_buffer_range,
5803 &mut data.rx_done,
5804 )
5805 .ok_or(WorkerError::InvalidRndisPacketCompletion)?;
5806 queue_state
5807 .queue
5808 .rx_avail(&mut queue_state.pool, &data.rx_done);
5809 }
5810 PacketData::SubChannelRequest(request) if state.primary.is_some() => {
5811 let mut subchannel_count = 0;
5812 let num_queues = request.num_sub_channels + 1;
5818 let status = if request.operation == protocol::SubchannelOperation::ALLOCATE
5819 && request.num_sub_channels < self.adapter.max_queues.into()
5820 && RxBufferRanges::validate_params(buffers.recv_buffer.count, num_queues)
5821 .is_ok()
5822 {
5823 subchannel_count = request.num_sub_channels;
5824 protocol::Status::SUCCESS
5825 } else {
5826 tracelimit::warn_ratelimited!(
5827 operation = ?request.operation,
5828 request_sub_channels = request.num_sub_channels,
5829 max_supported_sub_channels = self.adapter.max_queues - 1,
5830 recv_buffer_count = buffers.recv_buffer.count,
5831 "Subchannel request failed: either operation is not supported or requested more subchannels than supported"
5832 );
5833 protocol::Status::FAILURE
5834 };
5835
5836 tracing::debug!(?status, subchannel_count, "subchannel request");
5837 self.send_completion(
5838 packet.transaction_id,
5839 Some(&self.message(
5840 protocol::MESSAGE5_TYPE_SUB_CHANNEL,
5841 protocol::Message5SubchannelComplete {
5842 status,
5843 num_sub_channels: subchannel_count,
5844 },
5845 )),
5846 )?;
5847
5848 if subchannel_count > 0 {
5849 let primary = state.primary.as_mut().unwrap();
5850 primary.requested_num_queues = subchannel_count as u16 + 1;
5851 primary.tx_spread_sent = false;
5852 self.restart = Some(CoordinatorMessage::Restart { channel_idx: 0 });
5853 }
5854 }
5855 PacketData::RevokeReceiveBuffer(protocol::Message1RevokeReceiveBuffer { id })
5856 | PacketData::RevokeSendBuffer(protocol::Message1RevokeSendBuffer { id })
5857 if state.primary.is_some() =>
5858 {
5859 tracing::debug!(
5860 id,
5861 "receive/send buffer revoked, terminating channel processing"
5862 );
5863 return Err(WorkerError::BufferRevoked);
5864 }
5865 PacketData::SendVfAssociationCompletion if state.primary.is_some() => (),
5867 PacketData::SwitchDataPath(switch_data_path) if state.primary.is_some() => {
5868 self.switch_data_path(
5869 state,
5870 switch_data_path.active_data_path == protocol::DataPath::VF.0,
5871 packet.transaction_id,
5872 )?;
5873 }
5874 PacketData::SwitchDataPathCompletion if state.primary.is_some() => (),
5875 PacketData::OidQueryEx(oid_query) => {
5876 tracing::warn!(oid = ?oid_query.oid, "unimplemented OID");
5877 self.send_completion(
5878 packet.transaction_id,
5879 Some(&self.message(
5880 protocol::MESSAGE5_TYPE_OID_QUERY_EX_COMPLETE,
5881 protocol::Message5OidQueryExComplete {
5882 status: rndisprot::STATUS_NOT_SUPPORTED,
5883 bytes: 0,
5884 },
5885 )),
5886 )?;
5887 }
5888 p => {
5889 tracing::warn!(request = ?p, "unexpected packet");
5890 return Err(WorkerError::UnexpectedPacketOrder(
5891 PacketOrderError::SwitchDataPathCompletionPrimaryChannelState,
5892 ));
5893 }
5894 }
5895 }
5896 if total_packets > 0 && !self.transmit_segments(state, data, queue_state)? {
5897 state.stats.tx_stalled.increment();
5898 }
5899 state.stats.tx_packets_per_wake.add_sample(total_packets);
5900 Ok(did_some_work)
5901 }
5902
5903 fn transmit_pending_segments(
5906 &mut self,
5907 state: &mut ActiveState,
5908 data: &mut ProcessingData,
5909 queue_state: &mut QueueState,
5910 ) -> Result<bool, WorkerError> {
5911 if data.tx_segments.is_empty() {
5912 return Ok(false);
5913 }
5914 let sent = data.tx_segments_sent;
5915 let did_work =
5916 self.transmit_segments(state, data, queue_state)? || data.tx_segments_sent > sent;
5917 Ok(did_work)
5918 }
5919
5920 fn transmit_segments(
5922 &mut self,
5923 state: &mut ActiveState,
5924 data: &mut ProcessingData,
5925 queue_state: &mut QueueState,
5926 ) -> Result<bool, WorkerError> {
5927 let segments = &data.tx_segments[data.tx_segments_sent..];
5928 let (sync, segments_sent) = queue_state
5929 .queue
5930 .tx_avail(&mut queue_state.pool, segments)
5931 .map_err(WorkerError::Endpoint)?;
5932
5933 let mut segments = &segments[..segments_sent];
5934 data.tx_segments_sent += segments_sent;
5935
5936 if sync {
5937 while let Some(head) = segments.first() {
5939 let net_backend::TxSegmentType::Head(metadata) = &head.ty else {
5940 unreachable!()
5941 };
5942 let id = metadata.id;
5943 let pending_tx_packet = &mut state.pending_tx_packets[id.0 as usize];
5944 pending_tx_packet.pending_packet_count -= 1;
5945 if pending_tx_packet.pending_packet_count == 0 {
5946 self.complete_tx_packet(state, id, protocol::Status::SUCCESS)?;
5947 }
5948 segments = &segments[metadata.segment_count as usize..];
5949 }
5950 }
5951
5952 let all_sent = data.tx_segments_sent == data.tx_segments.len();
5953 if all_sent {
5954 data.tx_segments.clear();
5955 data.tx_segments_sent = 0;
5956 }
5957 Ok(all_sent)
5958 }
5959
5960 fn handle_rndis(
5961 &mut self,
5962 buffers: &ChannelBuffers,
5963 id: TxId,
5964 state: &mut ActiveState,
5965 packet: &Packet<'_>,
5966 segments: &mut Vec<TxSegment>,
5967 ) -> Result<usize, WorkerError> {
5968 let mut total_packets = 0;
5969 let tx_packet = &mut state.pending_tx_packets[id.0 as usize];
5970 assert_eq!(tx_packet.pending_packet_count, 0);
5971 tx_packet.transaction_id = packet
5972 .transaction_id
5973 .ok_or(WorkerError::MissingTransactionId)?;
5974
5975 packet
5979 .external_data
5980 .iter()
5981 .try_for_each(|range| buffers.mem.probe_gpns(range.gpns()))
5982 .map_err(WorkerError::GpaDirectError)?;
5983
5984 let mut reader = packet.rndis_reader(&buffers.mem);
5985 let header: rndisprot::MessageHeader = reader.read_plain()?;
5986 if header.message_type == rndisprot::MESSAGE_TYPE_PACKET_MSG {
5987 let start = segments.len();
5988 match self.handle_rndis_packet_messages(
5989 buffers,
5990 state,
5991 id,
5992 header.message_length as usize,
5993 reader,
5994 segments,
5995 ) {
5996 Ok(n) => {
5997 state.pending_tx_packets[id.0 as usize].pending_packet_count += n;
5998 total_packets += n;
5999 }
6000 Err(err) => {
6001 segments.truncate(start);
6003 return Err(err);
6004 }
6005 }
6006 } else {
6007 self.handle_rndis_message(state, header.message_type, reader)?;
6008 }
6009
6010 Ok(total_packets)
6011 }
6012
6013 fn try_send_tx_packet(
6014 &mut self,
6015 transaction_id: u64,
6016 status: protocol::Status,
6017 ) -> Result<bool, WorkerError> {
6018 let message = self.message(
6019 protocol::MESSAGE1_TYPE_SEND_RNDIS_PACKET_COMPLETE,
6020 protocol::Message1SendRndisPacketComplete { status },
6021 );
6022 let result = self.queue.split().1.batched().try_write_aligned(
6023 transaction_id,
6024 OutgoingPacketType::Completion,
6025 message.aligned_payload(),
6026 );
6027 let sent = match result {
6028 Ok(()) => true,
6029 Err(queue::TryWriteError::Full(n)) => {
6030 self.pending_send_size = n;
6031 false
6032 }
6033 Err(queue::TryWriteError::Queue(err)) => return Err(err.into()),
6034 };
6035 Ok(sent)
6036 }
6037
6038 fn send_pending_packets(&mut self, state: &mut ActiveState) -> Result<bool, WorkerError> {
6039 let mut did_some_work = false;
6040 while let Some(pending) = state.pending_tx_completions.front() {
6041 if !self.try_send_tx_packet(pending.transaction_id, pending.status)? {
6042 return Ok(did_some_work);
6043 }
6044 did_some_work = true;
6045 if let Some(id) = pending.tx_id {
6046 state.free_tx_packets.push(id);
6047 }
6048 tracing::trace!(?pending, "sent tx completion");
6049 state.pending_tx_completions.pop_front();
6050 }
6051
6052 self.pending_send_size = 0;
6053 Ok(did_some_work)
6054 }
6055
6056 fn complete_tx_packet(
6057 &mut self,
6058 state: &mut ActiveState,
6059 id: TxId,
6060 status: protocol::Status,
6061 ) -> Result<(), WorkerError> {
6062 let tx_packet = &mut state.pending_tx_packets[id.0 as usize];
6063 assert_eq!(tx_packet.pending_packet_count, 0);
6064 if self.pending_send_size == 0
6065 && self.try_send_tx_packet(tx_packet.transaction_id, status)?
6066 {
6067 tracing::trace!(id = id.0, "sent tx completion");
6068 state.free_tx_packets.push(id);
6069 } else {
6070 tracing::trace!(id = id.0, "pended tx completion");
6071 state.pending_tx_completions.push_back(PendingTxCompletion {
6072 transaction_id: tx_packet.transaction_id,
6073 tx_id: Some(id),
6074 status,
6075 });
6076 }
6077 Ok(())
6078 }
6079}
6080
6081impl ActiveState {
6082 fn release_recv_buffers(
6083 &mut self,
6084 transaction_id: u64,
6085 rx_buffer_range: &RxBufferRange,
6086 done: &mut Vec<RxId>,
6087 ) -> Option<()> {
6088 let first_id: u32 = transaction_id.try_into().ok()?;
6090 let ids = self.rx_bufs.free(first_id)?;
6091 for id in ids {
6092 if !rx_buffer_range.send_if_remote(id) {
6093 if id >= RX_RESERVED_CONTROL_BUFFERS {
6094 done.push(RxId(id));
6095 } else {
6096 self.primary
6097 .as_mut()?
6098 .free_control_buffers
6099 .push(ControlMessageId(id));
6100 }
6101 }
6102 }
6103 Some(())
6104 }
6105}