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