Skip to main content

netvsp/
lib.rs

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