Skip to main content

vmbus_server/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4#![expect(missing_docs)]
5#![forbid(unsafe_code)]
6
7mod channel_bitmap;
8pub mod channels;
9pub mod hvsock;
10mod monitor;
11mod proxyintegration;
12#[cfg(test)]
13mod tests;
14
15/// The GUID type used for vmbus channel identifiers.
16pub type Guid = guid::Guid;
17
18use anyhow::Context;
19use async_trait::async_trait;
20use channel_bitmap::ChannelBitmap;
21use channels::ConnectionTarget;
22pub use channels::InitiateContactRequest;
23use channels::MessageTarget;
24pub use channels::MnfUsage;
25use channels::ModifyConnectionRequest;
26use channels::ModifyConnectionResponse;
27use channels::Notifier;
28use channels::OfferId;
29pub use channels::OfferParamsInternal;
30use channels::OpenParams;
31use channels::RestoreError;
32pub use channels::Update;
33use futures::FutureExt;
34use futures::StreamExt;
35use futures::channel::mpsc;
36use futures::channel::mpsc::SendError;
37use futures::future::OptionFuture;
38use futures::future::poll_fn;
39use futures::stream::SelectAll;
40use guestmem::GuestMemory;
41use hvdef::Vtl;
42use inspect::Inspect;
43use mesh::payload::Protobuf;
44use mesh::rpc::FailableRpc;
45use mesh::rpc::Rpc;
46use mesh::rpc::RpcError;
47use mesh::rpc::RpcSend;
48use pal_async::driver::Driver;
49use pal_async::driver::SpawnDriver;
50use pal_async::task::Task;
51use pal_async::timer::PolledTimer;
52use pal_event::Event;
53#[cfg(windows)]
54pub use proxyintegration::ProxyIntegration;
55#[cfg(windows)]
56pub use proxyintegration::ProxyServerInfo;
57use ring::PAGE_SIZE;
58use std::collections::HashMap;
59use std::future;
60use std::future::Future;
61use std::pin::Pin;
62use std::sync::Arc;
63use std::task::Poll;
64use std::task::ready;
65use std::time::Duration;
66use unicycle::FuturesUnordered;
67use vmbus_channel::bus::ChannelRequest;
68use vmbus_channel::bus::ChannelServerRequest;
69use vmbus_channel::bus::GpadlRequest;
70use vmbus_channel::bus::ModifyRequest;
71use vmbus_channel::bus::OfferInput;
72use vmbus_channel::bus::OfferKey;
73use vmbus_channel::bus::OfferResources;
74use vmbus_channel::bus::OpenData;
75use vmbus_channel::bus::OpenRequest;
76use vmbus_channel::bus::ParentBus;
77use vmbus_channel::bus::RestoreResult;
78use vmbus_channel::gpadl::GpadlMap;
79use vmbus_channel::gpadl_ring::AlignedGpadlView;
80use vmbus_core::HvsockConnectRequest;
81use vmbus_core::HvsockConnectResult;
82use vmbus_core::MaxVersionInfo;
83use vmbus_core::OutgoingMessage;
84use vmbus_core::TaggedStream;
85use vmbus_core::VMBUS_SINT;
86use vmbus_core::VersionInfo;
87use vmbus_core::protocol;
88pub use vmbus_core::protocol::GpadlId;
89#[cfg(windows)]
90use vmbus_proxy::ProxyHandle;
91use vmbus_ring as ring;
92use vmbus_ring::gparange::MultiPagedRangeBuf;
93use vmcore::interrupt::Interrupt;
94use vmcore::save_restore::SavedStateRoot;
95use vmcore::synic::EventPort;
96use vmcore::synic::GuestEventPort;
97use vmcore::synic::GuestMessagePort;
98use vmcore::synic::MessagePort;
99use vmcore::synic::MonitorPageGpas;
100use vmcore::synic::SynicPortAccess;
101
102pub const REDIRECT_SINT: u8 = 7;
103pub const REDIRECT_VTL: Vtl = Vtl::Vtl2;
104const SHARED_EVENT_CONNECTION_ID: u32 = 2;
105const EVENT_PORT_ID: u32 = 2;
106const VMBUS_MESSAGE_TYPE: u32 = 1;
107
108const MAX_CONCURRENT_HVSOCK_REQUESTS: usize = 16;
109
110#[derive(Inspect)]
111pub struct VmbusServer {
112    #[inspect(flatten, send = "VmbusRequest::Inspect")]
113    task_send: mesh::Sender<VmbusRequest>,
114    #[inspect(skip)]
115    control: Arc<VmbusServerControl>,
116    #[inspect(skip)]
117    _message_port: Box<dyn Sync + Send>,
118    #[inspect(skip)]
119    _multiclient_message_port: Option<Box<dyn Sync + Send>>,
120    #[inspect(skip)]
121    task: Task<ServerTask>,
122}
123
124pub struct VmbusServerBuilder<T: SpawnDriver> {
125    spawner: T,
126    synic: Arc<dyn SynicPortAccess>,
127    gm: GuestMemory,
128    private_gm: Option<GuestMemory>,
129    vtl: Vtl,
130    hvsock_notify: Option<HvsockServerChannelHalf>,
131    server_relay: Option<VmbusServerChannelHalf>,
132    saved_state_notify: Option<mesh::Sender<SavedStateRequest>>,
133    external_server: Option<mesh::Sender<InitiateContactRequest>>,
134    external_requests: Option<mesh::Receiver<InitiateContactRequest>>,
135    use_message_redirect: bool,
136    channel_id_offset: u16,
137    max_version: Option<MaxVersionInfo>,
138    delay_max_version: bool,
139    max_restore_version: Option<MaxVersionInfo>,
140    enable_mnf: bool,
141    force_confidential_external_memory: bool,
142    support_gpa_pinning: bool,
143    force_gpa_pinning: bool,
144    send_messages_while_stopped: bool,
145    channel_unstick_delay: Option<Duration>,
146    use_absolute_channel_order: bool,
147}
148
149#[derive(mesh::MeshPayload)]
150/// The request to send to the proxy to set or clear its saved state cache.
151pub enum SavedStateRequest {
152    Set(FailableRpc<Box<channels::SavedState>, ()>),
153    Clear(Rpc<(), ()>),
154}
155
156/// The server side of the connection between a vmbus server and a relay.
157pub struct ServerChannelHalf<Request, Response> {
158    request_send: mesh::Sender<Request>,
159    response_receive: mesh::Receiver<Response>,
160}
161
162/// The relay side of a connection between a vmbus server and a relay.
163pub struct RelayChannelHalf<Request, Response> {
164    pub request_receive: mesh::Receiver<Request>,
165    pub response_send: mesh::Sender<Response>,
166}
167
168/// A connection between a vmbus server and a relay.
169pub struct RelayChannel<Request, Response> {
170    pub relay_half: RelayChannelHalf<Request, Response>,
171    pub server_half: ServerChannelHalf<Request, Response>,
172}
173
174impl<Request: 'static + Send, Response: 'static + Send> RelayChannel<Request, Response> {
175    /// Creates a new channel between the vmbus server and a relay.
176    pub fn new() -> Self {
177        let (request_send, request_receive) = mesh::channel();
178        let (response_send, response_receive) = mesh::channel();
179        Self {
180            relay_half: RelayChannelHalf {
181                request_receive,
182                response_send,
183            },
184            server_half: ServerChannelHalf {
185                request_send,
186                response_receive,
187            },
188        }
189    }
190}
191
192pub type VmbusServerChannelHalf = ServerChannelHalf<ModifyRelayRequest, ModifyRelayResponse>;
193pub type VmbusRelayChannelHalf = RelayChannelHalf<ModifyRelayRequest, ModifyRelayResponse>;
194pub type VmbusRelayChannel = RelayChannel<ModifyRelayRequest, ModifyRelayResponse>;
195pub type HvsockServerChannelHalf = ServerChannelHalf<HvsockConnectRequest, HvsockConnectResult>;
196pub type HvsockRelayChannelHalf = RelayChannelHalf<HvsockConnectRequest, HvsockConnectResult>;
197pub type HvsockRelayChannel = RelayChannel<HvsockConnectRequest, HvsockConnectResult>;
198
199/// A request from the server to the relay to modify connection state.
200///
201/// The version and use_interrupt_page fields can only be present if this request was sent for an
202/// InitiateContact message from the guest.
203///
204/// If `version` is `Some`, the relay must respond with either `ModifyRelayResponse::Supported` or
205/// `ModifyRelayResponse::Unsupported`. If `version` is `None`, the relay must respond with
206/// `ModifyRelayResponse::Modified`.
207#[derive(Debug, Copy, Clone)]
208pub struct ModifyRelayRequest {
209    pub version: Option<u32>,
210    pub monitor_page: Update<MonitorPageGpas>,
211    pub use_interrupt_page: Option<bool>,
212}
213
214/// A response from the relay to a ModifyRelayRequest from the server.
215#[derive(Debug, Copy, Clone)]
216pub enum ModifyRelayResponse {
217    /// The requested version change is supported, and the relay completed the connection
218    /// modification with the specified status. All of the feature flags supported by the relay host
219    /// are included, regardless of what features were requested.
220    Supported(protocol::ConnectionState, protocol::FeatureFlags),
221    /// A version change was requested but the relay host doesn't support that version. This
222    /// response cannot be returned for a request with no version change set.
223    Unsupported,
224    /// The connection modification completed with the specified status. This response must be sent
225    /// if no version change was requested.
226    Modified(protocol::ConnectionState),
227}
228
229impl From<ModifyConnectionRequest> for ModifyRelayRequest {
230    fn from(value: ModifyConnectionRequest) -> Self {
231        Self {
232            version: value.version.map(|v| v.version as u32),
233            monitor_page: value.monitor_page,
234            use_interrupt_page: match value.interrupt_page {
235                Update::Unchanged => None,
236                Update::Reset => Some(false),
237                Update::Set(_) => Some(true),
238            },
239        }
240    }
241}
242
243#[derive(Debug)]
244enum VmbusRequest {
245    Reset(Rpc<(), ()>),
246    Inspect(inspect::Deferred),
247    Save(Rpc<(), SavedState>),
248    Restore(Rpc<Box<SavedState>, Result<(), RestoreError>>),
249    Start,
250    Stop(Rpc<(), ()>),
251}
252
253#[derive(mesh::MeshPayload, Debug)]
254pub struct OfferInfo {
255    pub params: OfferParamsInternal,
256    pub event: Interrupt,
257    pub request_send: mesh::Sender<ChannelRequest>,
258    pub server_request_recv: mesh::Receiver<ChannelServerRequest>,
259}
260
261#[expect(clippy::large_enum_variant)]
262#[derive(mesh::MeshPayload)]
263pub(crate) enum OfferRequest {
264    Offer(FailableRpc<OfferInfo, ()>),
265    ForceReset(Rpc<(), ()>),
266}
267
268struct ChannelEvent(Interrupt);
269
270impl EventPort for ChannelEvent {
271    fn handle_event(&self, _flag: u16) {
272        self.0.deliver();
273    }
274
275    fn os_event(&self) -> Option<&Event> {
276        self.0.event()
277    }
278}
279
280#[derive(Debug, Protobuf, SavedStateRoot)]
281#[mesh(package = "vmbus.server")]
282pub struct SavedState {
283    #[mesh(1)]
284    pub server: channels::SavedState,
285    // Indicates if the lost synic bug is fixed or not. By default it's false.
286    // During the restore process, we check if the field is not true then
287    // unstick_channels() function will be called to mitigate the issue.
288    #[mesh(2)]
289    pub lost_synic_bug_fixed: bool,
290}
291
292const MESSAGE_CONNECTION_ID: u32 = 1;
293const MULTICLIENT_MESSAGE_CONNECTION_ID: u32 = 4;
294
295impl<T: SpawnDriver + Clone> VmbusServerBuilder<T> {
296    /// Creates a new builder for `VmbusServer` with the default options.
297    pub fn new(spawner: T, synic: Arc<dyn SynicPortAccess>, gm: GuestMemory) -> Self {
298        Self {
299            spawner,
300            synic,
301            gm,
302            private_gm: None,
303            vtl: Vtl::Vtl0,
304            hvsock_notify: None,
305            server_relay: None,
306            saved_state_notify: None,
307            external_server: None,
308            external_requests: None,
309            use_message_redirect: false,
310            channel_id_offset: 0,
311            max_version: None,
312            delay_max_version: false,
313            max_restore_version: None,
314            enable_mnf: false,
315            force_confidential_external_memory: false,
316            support_gpa_pinning: false,
317            force_gpa_pinning: false,
318            send_messages_while_stopped: false,
319            channel_unstick_delay: Some(Duration::from_millis(100)),
320            use_absolute_channel_order: false,
321        }
322    }
323
324    /// Sets a separate guest memory instance to use for channels that are confidential (non-relay
325    /// channels in Underhill on a hardware isolated VM). This is not relevant for a non-Underhill
326    /// VmBus server.
327    pub fn private_gm(mut self, private_gm: Option<GuestMemory>) -> Self {
328        self.private_gm = private_gm;
329        self
330    }
331
332    /// Sets the VTL that this instance will serve.
333    pub fn vtl(mut self, vtl: Vtl) -> Self {
334        self.vtl = vtl;
335        self
336    }
337
338    /// Sets a send/receive pair used to handle hvsocket requests.
339    pub fn hvsock_notify(mut self, hvsock_notify: Option<HvsockServerChannelHalf>) -> Self {
340        self.hvsock_notify = hvsock_notify;
341        self
342    }
343
344    /// Sets a send channel used to enlighten ProxyIntegration about saved channels.
345    pub fn saved_state_notify(
346        mut self,
347        saved_state_notify: Option<mesh::Sender<SavedStateRequest>>,
348    ) -> Self {
349        self.saved_state_notify = saved_state_notify;
350        self
351    }
352
353    /// Sets a send/receive pair that will be notified of server requests. This is used by the
354    /// Underhill relay.
355    pub fn server_relay(mut self, server_relay: Option<VmbusServerChannelHalf>) -> Self {
356        self.server_relay = server_relay;
357        self
358    }
359
360    /// Sets a receiver that receives requests from another server.
361    pub fn external_requests(
362        mut self,
363        external_requests: Option<mesh::Receiver<InitiateContactRequest>>,
364    ) -> Self {
365        self.external_requests = external_requests;
366        self
367    }
368
369    /// Sets a sender used to forward unhandled connect requests (which used a different VTL)
370    /// to another server.
371    pub fn external_server(
372        mut self,
373        external_server: Option<mesh::Sender<InitiateContactRequest>>,
374    ) -> Self {
375        self.external_server = external_server;
376        self
377    }
378
379    /// Sets a value which indicates whether the vmbus control plane is redirected to Underhill.
380    pub fn use_message_redirect(mut self, use_message_redirect: bool) -> Self {
381        self.use_message_redirect = use_message_redirect;
382        self
383    }
384
385    /// Tells the server to use an offset when generating channel IDs to void collisions with
386    /// another vmbus server.
387    ///
388    /// N.B. This should only be used by the Underhill vmbus server.
389    pub fn enable_channel_id_offset(mut self, enable: bool) -> Self {
390        self.channel_id_offset = if enable { 1024 } else { 0 };
391        self
392    }
393
394    /// Tells the server to limit the protocol version offered to the guest.
395    ///
396    /// N.B. This is used for testing older protocols without requiring a specific guest OS.
397    pub fn max_version(mut self, max_version: Option<MaxVersionInfo>) -> Self {
398        self.max_version = max_version;
399        self
400    }
401
402    /// Delay limiting the maximum version until after the first `Unload` message.
403    ///
404    /// N.B. This is used to enable the use of versions older than `Version::Win10` with Uefi boot,
405    ///      since that's the oldest version the Uefi client supports.
406    pub fn delay_max_version(mut self, delay: bool) -> Self {
407        self.delay_max_version = delay;
408        self
409    }
410
411    /// Tells the server to limit the protocol version accepted when restoring from saved state.
412    ///
413    /// This is configured separately from [`Self::max_version`] so that the version limit enforced
414    /// during restore can differ from the limit used for new connections. This allows a feature to
415    /// be available for rollback from a future version that has enabled it by default.
416    pub fn max_restore_version(mut self, max_restore_version: Option<MaxVersionInfo>) -> Self {
417        self.max_restore_version = max_restore_version;
418        self
419    }
420
421    /// Enable MNF support in the server.
422    ///
423    /// N.B. Enabling this has no effect if the synic does not support mapping monitor pages.
424    pub fn enable_mnf(mut self, enable: bool) -> Self {
425        self.enable_mnf = enable;
426        self
427    }
428
429    /// Force all non-relay channels to use encrypted external memory. Used for testing purposes
430    /// only.
431    pub fn force_confidential_external_memory(mut self, force: bool) -> Self {
432        self.force_confidential_external_memory = force;
433        self
434    }
435
436    /// Indicates the current VM environment supports GPA pinning so the relevant feature flag can
437    /// be used.
438    pub fn support_gpa_pinning(mut self, support: bool) -> Self {
439        self.support_gpa_pinning = support;
440        self
441    }
442
443    /// Force all channels to use pinned GPA ranges. Used for testing purposes only.
444    pub fn force_gpa_pinning(mut self, force: bool) -> Self {
445        self.force_gpa_pinning = force;
446        self
447    }
448
449    /// Send messages to the partition even while stopped, which can cause
450    /// corrupted synic states across VM reset.
451    ///
452    /// This option is used to prevent messages from getting into the queue, for
453    /// saved state compatibility with release/2411. It can be removed once that
454    /// release is no longer supported.
455    pub fn send_messages_while_stopped(mut self, send: bool) -> Self {
456        self.send_messages_while_stopped = send;
457        self
458    }
459
460    /// Sets the delay before unsticking a vmbus channel after it has been opened.
461    ///
462    /// This option provides a work around for guests that ignore interrupts before they receive the
463    /// OpenResult message, by triggering an interrupt after the channel has been opened.
464    ///
465    /// If not set, the default is 100ms. If set to `None`, no interrupt will be triggered.
466    pub fn channel_unstick_delay(mut self, delay: Option<Duration>) -> Self {
467        self.channel_unstick_delay = delay;
468        self
469    }
470
471    /// Sets whether the channel order value provided in an offer is the primary way of ordering
472    /// channels when assigning channel IDs, rather than the default behavior of ordering by
473    /// interface ID first.
474    pub fn use_absolute_channel_order(mut self, assign: bool) -> Self {
475        self.use_absolute_channel_order = assign;
476        self
477    }
478
479    /// Creates a new instance of the server.
480    ///
481    /// When the object is dropped, all channels will be closed and revoked
482    /// automatically.
483    pub fn build(self) -> anyhow::Result<VmbusServer> {
484        #[expect(clippy::disallowed_methods)] // TODO
485        let (message_send, message_recv) = mpsc::channel(64);
486        let message_sender = Arc::new(MessageSender {
487            send: message_send.clone(),
488            multiclient: self.use_message_redirect,
489        });
490
491        let (redirect_vtl, redirect_sint) = if self.use_message_redirect {
492            (REDIRECT_VTL, REDIRECT_SINT)
493        } else {
494            (self.vtl, VMBUS_SINT)
495        };
496
497        // If this server is not for VTL2, use a server-specific connection ID rather than the
498        // standard one.
499        let connection_id = if self.vtl == Vtl::Vtl0 && !self.use_message_redirect {
500            MESSAGE_CONNECTION_ID
501        } else {
502            // TODO: This ID should be using the correct target VP, but that is not known until
503            //       InitiateContact.
504            VmbusServer::get_child_message_connection_id(0, redirect_sint, redirect_vtl)
505        };
506
507        let _message_port = self
508            .synic
509            .add_message_port(connection_id, redirect_vtl, message_sender)
510            .context("failed to create vmbus synic ports")?;
511
512        // If this server is for VTL0, it is also responsible for the multiclient message port.
513        // N.B. If control plane redirection is enabled, the redirected message port is used for
514        //      multiclient and no separate multiclient port is created.
515        let _multiclient_message_port = if self.vtl == Vtl::Vtl0 && !self.use_message_redirect {
516            let multiclient_message_sender = Arc::new(MessageSender {
517                send: message_send,
518                multiclient: true,
519            });
520
521            Some(
522                self.synic
523                    .add_message_port(
524                        MULTICLIENT_MESSAGE_CONNECTION_ID,
525                        self.vtl,
526                        multiclient_message_sender,
527                    )
528                    .context("failed to create vmbus synic ports")?,
529            )
530        } else {
531            None
532        };
533
534        let (offer_send, offer_recv) = mesh::mpsc_channel();
535        let control = Arc::new(VmbusServerControl {
536            mem: self.gm.clone(),
537            private_mem: self.private_gm.clone(),
538            send: offer_send,
539            use_event: self.synic.prefer_os_events(),
540            force_confidential_external_memory: self.force_confidential_external_memory,
541            force_gpa_pinning: self.force_gpa_pinning,
542        });
543
544        let mut server = channels::Server::new(
545            self.vtl,
546            connection_id,
547            self.channel_id_offset,
548            self.use_absolute_channel_order,
549            self.support_gpa_pinning,
550        );
551
552        // If MNF is handled by this server and this is a paravisor for an isolated VM, the monitor
553        // pages must be allocated by the server, not the guest, since the guest will provide shared
554        // pages which can't be used in this case. If the guest doesn't support server-specified
555        // monitor pages, MNF will be disabled for all channels for that connection.
556        server.set_require_server_allocated_mnf(self.enable_mnf && self.private_gm.is_some());
557
558        // If requested, limit the maximum protocol version and feature flags.
559        if let Some(version) = self.max_version {
560            server.set_compatibility_version(version, self.delay_max_version);
561        }
562
563        if let Some(version) = self.max_restore_version {
564            server.set_restore_compatibility_version(version);
565        }
566
567        let (relay_request_send, relay_response_recv) =
568            if let Some(server_relay) = self.server_relay {
569                let r = server_relay.response_receive.boxed().fuse();
570                (server_relay.request_send, r)
571            } else {
572                let (req_send, req_recv) = mesh::channel();
573                let resp_recv = req_recv
574                    .map(|req: ModifyRelayRequest| {
575                        // Map to the correct response type for the request.
576                        if req.version.is_some() {
577                            ModifyRelayResponse::Supported(
578                                protocol::ConnectionState::SUCCESSFUL,
579                                protocol::FeatureFlags::from_bits(u32::MAX),
580                            )
581                        } else {
582                            ModifyRelayResponse::Modified(protocol::ConnectionState::SUCCESSFUL)
583                        }
584                    })
585                    .boxed()
586                    .fuse();
587                (req_send, resp_recv)
588            };
589
590        // If no hvsock notifier was specified, use a default one that always sends an error response.
591        let (hvsock_send, hvsock_recv) = if let Some(hvsock_notify) = self.hvsock_notify {
592            let r = hvsock_notify.response_receive.boxed().fuse();
593            (hvsock_notify.request_send, r)
594        } else {
595            let (req_send, req_recv) = mesh::channel();
596            let resp_recv = req_recv
597                .map(|r: HvsockConnectRequest| HvsockConnectResult::from_request(&r, false))
598                .boxed()
599                .fuse();
600            (req_send, resp_recv)
601        };
602
603        let inner = ServerTaskInner {
604            running: false,
605            send_messages_while_stopped: self.send_messages_while_stopped,
606            gm: self.gm,
607            private_gm: self.private_gm,
608            vtl: self.vtl,
609            redirect_vtl,
610            redirect_sint,
611            message_port: self
612                .synic
613                .new_guest_message_port(redirect_vtl, 0, redirect_sint)?,
614            synic: self.synic,
615            hvsock_requests: 0,
616            hvsock_send,
617            saved_state_notify: self.saved_state_notify,
618            channels: HashMap::new(),
619            channel_responses: FuturesUnordered::new(),
620            relay_send: relay_request_send,
621            external_server_send: self.external_server,
622            channel_bitmap: None,
623            shared_event_port: None,
624            reset_done: Vec::new(),
625            mnf_support: self.enable_mnf.then(MnfSupport::default),
626        };
627
628        let (task_send, task_recv) = mesh::channel();
629        let mut server_task = ServerTask {
630            driver: Box::new(self.spawner.clone()),
631            server,
632            task_recv,
633            offer_recv,
634            message_recv,
635            server_request_recv: SelectAll::new(),
636            inner,
637            external_requests: self.external_requests,
638            next_seq: 0,
639            perform_post_restore_on_start: false,
640            unstick_on_start: false,
641            channel_unstickers: FuturesUnordered::new(),
642            channel_unstick_delay: self.channel_unstick_delay,
643        };
644
645        let task = self.spawner.spawn("vmbus server", async move {
646            server_task.run(relay_response_recv, hvsock_recv).await;
647            server_task
648        });
649
650        Ok(VmbusServer {
651            task_send,
652            control,
653            _message_port,
654            _multiclient_message_port,
655            task,
656        })
657    }
658}
659
660impl VmbusServer {
661    /// Creates a new builder for `VmbusServer` with the default options.
662    pub fn builder<T: SpawnDriver + Clone>(
663        spawner: T,
664        synic: Arc<dyn SynicPortAccess>,
665        gm: GuestMemory,
666    ) -> VmbusServerBuilder<T> {
667        VmbusServerBuilder::new(spawner, synic, gm)
668    }
669
670    pub async fn save(&self) -> SavedState {
671        self.task_send.call(VmbusRequest::Save, ()).await.unwrap()
672    }
673
674    pub async fn restore(&self, state: SavedState) -> Result<(), RestoreError> {
675        self.task_send
676            .call(VmbusRequest::Restore, Box::new(state))
677            .await
678            .unwrap()
679    }
680
681    /// Stop the control plane.
682    pub async fn stop(&self) {
683        self.task_send.call(VmbusRequest::Stop, ()).await.unwrap()
684    }
685
686    /// Starts the control plane.
687    pub fn start(&self) {
688        self.task_send.send(VmbusRequest::Start);
689    }
690
691    /// Resets the vmbus channel state.
692    pub async fn reset(&self) {
693        tracing::debug!("resetting channel state");
694        self.task_send.call(VmbusRequest::Reset, ()).await.unwrap()
695    }
696
697    /// Tears down the vmbus control plane.
698    pub async fn shutdown(self) {
699        drop(self.task_send);
700        let _ = self.task.await;
701    }
702
703    /// Returns an object that can be used to offer channels.
704    pub fn control(&self) -> Arc<VmbusServerControl> {
705        self.control.clone()
706    }
707
708    /// Returns the message connection ID to use for a communication from the guest for servers
709    /// that use a non-standard SINT or VTL.
710    fn get_child_message_connection_id(vp_index: u32, sint_index: u8, vtl: Vtl) -> u32 {
711        MULTICLIENT_MESSAGE_CONNECTION_ID
712            | (vtl as u32) << 22
713            | vp_index << 8
714            | (sint_index as u32) << 4
715    }
716
717    fn get_child_event_port_id(channel_id: protocol::ChannelId, sint_index: u8, vtl: Vtl) -> u32 {
718        EVENT_PORT_ID | (vtl as u32) << 22 | channel_id.0 << 8 | (sint_index as u32) << 4
719    }
720}
721
722#[derive(Default)]
723pub struct SynicMessage {
724    data: Vec<u8>,
725    multiclient: bool,
726    trusted: bool,
727}
728
729/// Information used by a server that supports MNF.
730#[derive(Default)]
731struct MnfSupport {
732    allocated_monitor_page: Option<MonitorPageGpas>,
733}
734
735/// Disambiguates offer instances that may have reused the same offer ID.
736#[derive(Debug, Clone, Copy)]
737struct OfferInstanceId {
738    offer_id: OfferId,
739    seq: u64,
740}
741
742struct ServerTask {
743    driver: Box<dyn Driver>,
744    server: channels::Server,
745    task_recv: mesh::Receiver<VmbusRequest>,
746    offer_recv: mesh::Receiver<OfferRequest>,
747    message_recv: mpsc::Receiver<SynicMessage>,
748    server_request_recv:
749        SelectAll<TaggedStream<OfferInstanceId, mesh::Receiver<ChannelServerRequest>>>,
750    inner: ServerTaskInner,
751    external_requests: Option<mesh::Receiver<InitiateContactRequest>>,
752    /// Next value for [`Channel::seq`].
753    next_seq: u64,
754    perform_post_restore_on_start: bool,
755    unstick_on_start: bool,
756    channel_unstickers: FuturesUnordered<Pin<Box<dyn Send + Future<Output = OfferInstanceId>>>>,
757    channel_unstick_delay: Option<Duration>,
758}
759
760struct ServerTaskInner {
761    running: bool,
762    send_messages_while_stopped: bool,
763    gm: GuestMemory,
764    private_gm: Option<GuestMemory>,
765    synic: Arc<dyn SynicPortAccess>,
766    vtl: Vtl,
767    redirect_vtl: Vtl,
768    redirect_sint: u8,
769    message_port: Box<dyn GuestMessagePort>,
770    hvsock_requests: usize,
771    hvsock_send: mesh::Sender<HvsockConnectRequest>,
772    saved_state_notify: Option<mesh::Sender<SavedStateRequest>>,
773    channels: HashMap<OfferId, Channel>,
774    channel_responses: FuturesUnordered<
775        Pin<Box<dyn Send + Future<Output = (OfferId, u64, Result<ChannelResponse, RpcError>)>>>,
776    >,
777    external_server_send: Option<mesh::Sender<InitiateContactRequest>>,
778    relay_send: mesh::Sender<ModifyRelayRequest>,
779    channel_bitmap: Option<Arc<ChannelBitmap>>,
780    shared_event_port: Option<Box<dyn Send>>,
781    reset_done: Vec<Rpc<(), ()>>,
782    /// Stores information needed to support MNF. If `None`, this server doesn't support MNF (in
783    /// the case of OpenHCL, that means it will be handled by the relay host).
784    mnf_support: Option<MnfSupport>,
785}
786
787#[derive(Debug)]
788enum ChannelResponse {
789    Open(bool),
790    Close,
791    Gpadl(GpadlId, bool),
792    TeardownGpadl(GpadlId),
793    Modify(i32),
794}
795
796#[derive(Debug, Copy, Clone, PartialEq, Eq)]
797enum ChannelUnstickState {
798    None,
799    Queued,
800    NeedsRequeue,
801}
802
803struct Channel {
804    key: OfferKey,
805    send: mesh::Sender<ChannelRequest>,
806    seq: u64,
807    state: ChannelState,
808    gpadls: Arc<GpadlMap>,
809    guest_to_host_event: Arc<ChannelEvent>,
810    flags: protocol::OfferFlags,
811    // A channel can be reserved no matter what state it is in. This allows the message port for a
812    // reserved channel to remain available even if the channel is closed, so the guest can read the
813    // close reserved channel response. The reserved state is cleared when the channel is revoked,
814    // reopened, or the guest sends an unload message.
815    reserved_state: ReservedState,
816    unstick_state: ChannelUnstickState,
817}
818
819struct ReservedState {
820    message_port: Option<Box<dyn GuestMessagePort>>,
821    target: ConnectionTarget,
822}
823
824struct ChannelOpenState {
825    open_params: OpenParams,
826    _event_port: Box<dyn Send>,
827    guest_event_port: Option<Box<dyn GuestEventPort>>,
828    host_to_guest_interrupt: Interrupt,
829}
830
831impl ChannelOpenState {
832    fn set_event_port_target_vp(&mut self, vp: u32) -> anyhow::Result<()> {
833        let Some(guest_event_port) = self.guest_event_port.as_mut() else {
834            anyhow::bail!("cannot set target VP if the channel interrupt is disabled");
835        };
836
837        guest_event_port.set_target_vp(vp)?;
838        Ok(())
839    }
840}
841
842enum ChannelState {
843    Closed,
844    Open(Box<ChannelOpenState>),
845    Closing,
846}
847
848impl ServerTask {
849    fn handle_offer(&mut self, mut info: OfferInfo) -> anyhow::Result<()> {
850        let key = info.params.key();
851        let flags = info.params.flags;
852
853        if self.inner.mnf_support.is_some() && self.inner.synic.monitor_support().is_some() {
854            // If this server is handling MnF, ignore any relayed monitor IDs but still enable MnF
855            // for those channels.
856            // N.B. Since this can only happen in OpenHCL, which emulates MnF, the latency is
857            //      ignored.
858            if info.params.use_mnf.is_relayed() {
859                info.params.use_mnf = MnfUsage::Enabled {
860                    latency: Duration::ZERO,
861                }
862            }
863        } else if info.params.use_mnf.is_enabled() {
864            // If the server is not handling MnF, disable it for the channel. This does not affect
865            // channels with a relayed monitor ID.
866            info.params.use_mnf = MnfUsage::Disabled;
867        }
868
869        let offer_id = self
870            .server
871            .with_notifier(&mut self.inner)
872            .offer_channel(info.params)
873            .context("channel offer failed")?;
874
875        tracing::debug!(?offer_id, %key, "offered channel");
876
877        let seq = self.next_seq;
878        self.next_seq += 1;
879        self.inner.channels.insert(
880            offer_id,
881            Channel {
882                key,
883                send: info.request_send,
884                state: ChannelState::Closed,
885                gpadls: GpadlMap::new(),
886                guest_to_host_event: Arc::new(ChannelEvent(info.event)),
887                seq,
888                flags,
889                reserved_state: ReservedState {
890                    message_port: None,
891                    target: ConnectionTarget { vp: 0, sint: 0 },
892                },
893                unstick_state: ChannelUnstickState::None,
894            },
895        );
896
897        self.server_request_recv.push(TaggedStream::new(
898            OfferInstanceId { offer_id, seq },
899            info.server_request_recv,
900        ));
901
902        Ok(())
903    }
904
905    fn handle_revoke(&mut self, id: OfferInstanceId) {
906        // The channel may or may not exist in the map depending on whether it's been explicitly
907        // revoked before being dropped.
908        if let Some(channel) = self.inner.channels.get(&id.offer_id) {
909            if channel.seq == id.seq {
910                tracing::info!(?id.offer_id, key = %channel.key, "revoking channel");
911                self.inner.channels.remove(&id.offer_id);
912                self.server
913                    .with_notifier(&mut self.inner)
914                    .revoke_channel(id.offer_id);
915            }
916        }
917    }
918
919    fn handle_response(
920        &mut self,
921        offer_id: OfferId,
922        seq: u64,
923        response: Result<ChannelResponse, RpcError>,
924    ) {
925        // Validate the sequence to ensure the response is not for a revoked channel.
926        let channel = self
927            .inner
928            .channels
929            .get(&offer_id)
930            .filter(|channel| channel.seq == seq);
931
932        if let Some(channel) = channel {
933            match response {
934                Ok(response) => match response {
935                    ChannelResponse::Open(result) => self.handle_open(offer_id, result),
936                    ChannelResponse::Close => self.handle_close(offer_id),
937                    ChannelResponse::Gpadl(gpadl_id, ok) => {
938                        self.handle_gpadl_create(offer_id, gpadl_id, ok)
939                    }
940                    ChannelResponse::TeardownGpadl(gpadl_id) => {
941                        self.handle_gpadl_teardown(offer_id, gpadl_id)
942                    }
943                    ChannelResponse::Modify(status) => self.handle_modify_channel(offer_id, status),
944                },
945                Err(err) => {
946                    tracing::error!(
947                        key = %channel.key,
948                        error = &err as &dyn std::error::Error,
949                        "channel response failure, channel is in inconsistent state until revoked"
950                    );
951                }
952            }
953        } else {
954            tracing::debug!(offer_id = ?offer_id, seq, ?response, "received response after revoke");
955        }
956    }
957
958    fn handle_open(&mut self, offer_id: OfferId, success: bool) {
959        let status = if success {
960            let channel = self
961                .inner
962                .channels
963                .get_mut(&offer_id)
964                .expect("channel exists");
965
966            // Some guests ignore interrupts before they receive the OpenResult message. To avoid
967            // a potential hang, signal the channel after a delay if needed.
968            if let Some(delay) = self.channel_unstick_delay {
969                if channel.unstick_state == ChannelUnstickState::None {
970                    channel.unstick_state = ChannelUnstickState::Queued;
971                    let seq = channel.seq;
972                    let mut timer = PolledTimer::new(&self.driver);
973                    self.channel_unstickers.push(Box::pin(async move {
974                        timer.sleep(delay).await;
975                        OfferInstanceId { offer_id, seq }
976                    }));
977                } else {
978                    channel.unstick_state = ChannelUnstickState::NeedsRequeue;
979                }
980            }
981
982            0
983        } else {
984            protocol::STATUS_UNSUCCESSFUL
985        };
986
987        self.server
988            .with_notifier(&mut self.inner)
989            .open_complete(offer_id, status);
990    }
991
992    fn handle_close(&mut self, offer_id: OfferId) {
993        let channel = self
994            .inner
995            .channels
996            .get_mut(&offer_id)
997            .expect("channel still exists");
998
999        match &mut channel.state {
1000            ChannelState::Closing => {
1001                tracing::debug!(?offer_id, key = %channel.key, "closing channel");
1002                channel.state = ChannelState::Closed;
1003                self.server
1004                    .with_notifier(&mut self.inner)
1005                    .close_complete(offer_id);
1006            }
1007            _ => {
1008                tracing::error!(?offer_id, key = %channel.key, "invalid close channel response");
1009            }
1010        };
1011    }
1012
1013    fn handle_gpadl_create(&mut self, offer_id: OfferId, gpadl_id: GpadlId, ok: bool) {
1014        let status = if ok { 0 } else { protocol::STATUS_UNSUCCESSFUL };
1015        self.server
1016            .with_notifier(&mut self.inner)
1017            .gpadl_create_complete(offer_id, gpadl_id, status);
1018    }
1019
1020    fn handle_gpadl_teardown(&mut self, offer_id: OfferId, gpadl_id: GpadlId) {
1021        self.server
1022            .with_notifier(&mut self.inner)
1023            .gpadl_teardown_complete(offer_id, gpadl_id);
1024    }
1025
1026    fn handle_modify_channel(&mut self, offer_id: OfferId, status: i32) {
1027        self.server
1028            .with_notifier(&mut self.inner)
1029            .modify_channel_complete(offer_id, status);
1030    }
1031
1032    fn handle_restore_channel(
1033        &mut self,
1034        offer_id: OfferId,
1035        open: bool,
1036    ) -> anyhow::Result<RestoreResult> {
1037        let gpadls = self.server.channel_gpadls(offer_id);
1038
1039        // If the channel is opened, handle that before calling into channels so that failure can
1040        // be handled before the channel is marked restored.
1041        let open_request = open
1042            .then(|| -> anyhow::Result<_> {
1043                let params = self.server.get_restore_open_params(offer_id)?;
1044                let (channel, interrupt) = self.inner.open_channel(offer_id, &params)?;
1045                Ok(OpenRequest::new(
1046                    params.open_data,
1047                    interrupt,
1048                    self.server
1049                        .get_version()
1050                        .expect("must be connected")
1051                        .feature_flags,
1052                    channel.flags,
1053                ))
1054            })
1055            .transpose()?;
1056
1057        self.server
1058            .with_notifier(&mut self.inner)
1059            .restore_channel(offer_id, open_request.is_some())?;
1060
1061        let channel = self.inner.channels.get_mut(&offer_id).unwrap();
1062        for gpadl in &gpadls {
1063            if let Ok(buf) = MultiPagedRangeBuf::from_range_buffer(
1064                gpadl.request.count.into(),
1065                gpadl.request.buf.clone(),
1066            ) {
1067                channel.gpadls.add(gpadl.request.id, buf);
1068            }
1069        }
1070
1071        let result = RestoreResult {
1072            open_request,
1073            gpadls,
1074        };
1075        Ok(result)
1076    }
1077
1078    async fn handle_request(&mut self, request: VmbusRequest) {
1079        tracing::debug!(?request, "handle_request");
1080        match request {
1081            VmbusRequest::Reset(rpc) => self.handle_reset(rpc),
1082            VmbusRequest::Inspect(deferred) => {
1083                deferred.respond(|resp| {
1084                    resp.field("message_port", &self.inner.message_port)
1085                        .field("running", self.inner.running)
1086                        .field("hvsock_requests", self.inner.hvsock_requests)
1087                        .field("channel_unstick_delay", self.channel_unstick_delay)
1088                        .field_mut_with("unstick_channels", |v| {
1089                            let v: inspect::ValueKind = if let Some(v) = v {
1090                                if v == "force" {
1091                                    self.unstick_channels(true);
1092                                    v.into()
1093                                } else {
1094                                    let v =
1095                                        v.parse().ok().context("expected false, true, or force")?;
1096                                    if v {
1097                                        self.unstick_channels(false);
1098                                    }
1099                                    v.into()
1100                                }
1101                            } else {
1102                                false.into()
1103                            };
1104                            anyhow::Ok(v)
1105                        })
1106                        .merge(&self.server.with_notifier(&mut self.inner));
1107                });
1108            }
1109            VmbusRequest::Save(rpc) => rpc.handle_sync(|()| SavedState {
1110                server: self.server.save(),
1111                lost_synic_bug_fixed: true,
1112            }),
1113            VmbusRequest::Restore(rpc) => {
1114                rpc.handle(async |state| {
1115                    self.unstick_on_start = !state.lost_synic_bug_fixed;
1116                    self.perform_post_restore_on_start = true;
1117                    if let Some(sender) = &self.inner.saved_state_notify {
1118                        tracing::trace!("sending saved state to proxy");
1119                        if let Err(err) = sender
1120                            .call_failable(SavedStateRequest::Set, Box::new(state.server.clone()))
1121                            .await
1122                        {
1123                            tracing::error!(
1124                                err = &err as &dyn std::error::Error,
1125                                "failed to restore proxy saved state"
1126                            );
1127                            return Err(RestoreError::ServerError(err.into()));
1128                        }
1129                    }
1130
1131                    self.server
1132                        .with_notifier(&mut self.inner)
1133                        .restore(state.server)
1134                })
1135                .await
1136            }
1137            VmbusRequest::Stop(rpc) => rpc.handle_sync(|()| {
1138                if self.inner.running {
1139                    self.inner.running = false;
1140                }
1141            }),
1142            VmbusRequest::Start => {
1143                if !self.inner.running {
1144                    self.inner.running = true;
1145                    if self.perform_post_restore_on_start {
1146                        if let Some(sender) = self.inner.saved_state_notify.as_ref() {
1147                            // Indicate to the proxy that the server is starting and that it should
1148                            // clear its saved state cache.
1149                            tracing::trace!("sending clear saved state message to proxy");
1150                            sender
1151                                .call(SavedStateRequest::Clear, ())
1152                                .await
1153                                .expect("failed to clear proxy saved state");
1154                        }
1155
1156                        self.server
1157                            .with_notifier(&mut self.inner)
1158                            .revoke_unclaimed_channels();
1159
1160                        self.perform_post_restore_on_start = false;
1161                    }
1162
1163                    if self.unstick_on_start {
1164                        tracing::info!(
1165                            "lost synic bug fix is not in yet, call unstick_channels to mitigate the issue."
1166                        );
1167                        self.unstick_channels(false);
1168                        self.unstick_on_start = false;
1169                    }
1170                }
1171            }
1172        }
1173    }
1174
1175    fn handle_reset(&mut self, rpc: Rpc<(), ()>) {
1176        let needs_reset = self.inner.reset_done.is_empty();
1177        self.inner.reset_done.push(rpc);
1178        if needs_reset {
1179            self.server.with_notifier(&mut self.inner).reset();
1180        }
1181    }
1182
1183    fn handle_relay_response(&mut self, response: ModifyRelayResponse) {
1184        // Convert to a matching ModifyConnectionResponse.
1185        let response = match response {
1186            ModifyRelayResponse::Supported(state, features) => {
1187                // Provide the server-allocated monitor page to the server only if they're actually being
1188                // used (if not, they may still be allocated from a previous connection).
1189                let allocated_monitor_gpas = self
1190                    .inner
1191                    .mnf_support
1192                    .as_ref()
1193                    .and_then(|mnf| mnf.allocated_monitor_page);
1194
1195                ModifyConnectionResponse::Supported(state, features, allocated_monitor_gpas)
1196            }
1197            ModifyRelayResponse::Unsupported => ModifyConnectionResponse::Unsupported,
1198            ModifyRelayResponse::Modified(state) => ModifyConnectionResponse::Modified(state),
1199        };
1200
1201        self.server
1202            .with_notifier(&mut self.inner)
1203            .complete_modify_connection(response);
1204    }
1205
1206    fn handle_tl_connect_result(&mut self, result: HvsockConnectResult) {
1207        tracing::debug!(?result, "hvsock connect result");
1208        assert_ne!(self.inner.hvsock_requests, 0);
1209        self.inner.hvsock_requests -= 1;
1210
1211        self.server
1212            .with_notifier(&mut self.inner)
1213            .send_tl_connect_result(result);
1214    }
1215
1216    fn handle_synic_message(&mut self, message: SynicMessage) {
1217        match self
1218            .server
1219            .with_notifier(&mut self.inner)
1220            .handle_synic_message(message)
1221        {
1222            Ok(()) => {}
1223            Err(err) => {
1224                tracing::warn!(
1225                    error = &err as &dyn std::error::Error,
1226                    "synic message error"
1227                );
1228            }
1229        }
1230    }
1231
1232    /// Handles a request forwarded by a different vmbus server. This is used to forward requests
1233    /// for different VTLs to different servers.
1234    ///
1235    /// N.B. This uses the same mechanism as the HCL server relay, so all requests, even the ones
1236    ///      meant for the primary server, are forwarded. In that case the primary server depends
1237    ///      on this server to send back a response so it can continue handling it.
1238    fn handle_external_request(&mut self, request: InitiateContactRequest) {
1239        self.server
1240            .with_notifier(&mut self.inner)
1241            .initiate_contact(request);
1242    }
1243
1244    async fn run(
1245        &mut self,
1246        mut relay_response_recv: impl futures::stream::FusedStream<Item = ModifyRelayResponse> + Unpin,
1247        mut hvsock_recv: impl futures::stream::FusedStream<Item = HvsockConnectResult> + Unpin,
1248    ) {
1249        loop {
1250            // Create an OptionFuture for each event that should only be handled
1251            // while the VM is running. In other cases, leave the events in
1252            // their respective queues.
1253
1254            let running_not_resetting = self.inner.running && self.inner.reset_done.is_empty();
1255            let mut external_requests = OptionFuture::from(
1256                running_not_resetting
1257                    .then(|| {
1258                        self.external_requests
1259                            .as_mut()
1260                            .map(|r| r.select_next_some())
1261                    })
1262                    .flatten(),
1263            );
1264
1265            // Try to send any pending messages while the VM is running.
1266            let has_pending_messages = self.server.has_pending_messages();
1267            let message_port = self.inner.message_port.as_mut();
1268            let mut flush_pending_messages =
1269                OptionFuture::from((running_not_resetting && has_pending_messages).then(|| {
1270                    poll_fn(|cx| {
1271                        self.server.poll_flush_pending_messages(|msg| {
1272                            message_port.poll_post_message(cx, VMBUS_MESSAGE_TYPE, msg.data())
1273                        })
1274                    })
1275                    .fuse()
1276                }));
1277
1278            // Only handle new incoming messages if there are no outgoing messages pending, and not
1279            // too many hvsock requests outstanding. This puts a bound on the resources used by the
1280            // guest.
1281            let mut message_recv = OptionFuture::from(
1282                (running_not_resetting
1283                    && !has_pending_messages
1284                    && self.inner.hvsock_requests < MAX_CONCURRENT_HVSOCK_REQUESTS)
1285                    .then(|| self.message_recv.select_next_some()),
1286            );
1287
1288            // Accept channel responses until stopped or when resetting.
1289            let mut channel_response = OptionFuture::from(
1290                (self.inner.running || !self.inner.reset_done.is_empty())
1291                    .then(|| self.inner.channel_responses.select_next_some()),
1292            );
1293
1294            // Accept hvsock connect responses while the VM is running.
1295            let mut hvsock_response =
1296                OptionFuture::from(running_not_resetting.then(|| hvsock_recv.select_next_some()));
1297
1298            let mut channel_unstickers = OptionFuture::from(
1299                running_not_resetting.then(|| self.channel_unstickers.select_next_some()),
1300            );
1301
1302            futures::select! { // merge semantics
1303                r = self.task_recv.recv().fuse() => {
1304                    if let Ok(request) = r {
1305                        self.handle_request(request).await;
1306                    } else {
1307                        break;
1308                    }
1309                }
1310                r = self.offer_recv.select_next_some() => {
1311                    match r {
1312                        OfferRequest::Offer(rpc) => {
1313                            rpc.handle_failable_sync(|request| { self.handle_offer(request) })
1314                        },
1315                        OfferRequest::ForceReset(rpc) => {
1316                            self.handle_reset(rpc);
1317                        }
1318                    }
1319                }
1320                r = self.server_request_recv.select_next_some() => {
1321                    match r {
1322                        (id, Some(request)) => match request {
1323                            ChannelServerRequest::Restore(rpc) => rpc.handle_failable_sync(|open| {
1324                                self.handle_restore_channel(id.offer_id, open)
1325                            }),
1326                            ChannelServerRequest::Revoke(rpc) => rpc.handle_sync(|_| {
1327                                self.handle_revoke(id);
1328                            })
1329                        },
1330                        (id, None) => self.handle_revoke(id),
1331                    }
1332                }
1333                r = channel_response => {
1334                    let (id, seq, response) = r.unwrap();
1335                    self.handle_response(id, seq, response);
1336                }
1337                r = relay_response_recv.select_next_some() => {
1338                    self.handle_relay_response(r);
1339                },
1340                r = hvsock_response => {
1341                    self.handle_tl_connect_result(r.unwrap());
1342                }
1343                data = message_recv => {
1344                    let data = data.unwrap();
1345                    self.handle_synic_message(data);
1346                }
1347                r = external_requests => {
1348                    let r = r.unwrap();
1349                    self.handle_external_request(r);
1350                }
1351                r = channel_unstickers => {
1352                    self.unstick_channel_by_id(r.unwrap());
1353                }
1354                _r = flush_pending_messages => {}
1355                complete => break,
1356            }
1357        }
1358    }
1359
1360    /// Wakes the guest and optionally the host for every open channel. If `force`, always wakes
1361    /// them. If `!force`, only wake for rings that are in the state where a notification is
1362    /// expected.
1363    fn unstick_channels(&self, force: bool) {
1364        let Some(version) = self.server.get_version() else {
1365            tracing::warn!("cannot unstick when not connected");
1366            return;
1367        };
1368
1369        for channel in self.inner.channels.values() {
1370            let gm = self.inner.get_gm_for_channel(version, channel);
1371            if let Err(err) = Self::unstick_channel(gm, channel, force, true) {
1372                tracing::warn!(
1373                    channel = %channel.key,
1374                    error = err.as_ref() as &dyn std::error::Error,
1375                    "could not unstick channel"
1376                );
1377            }
1378        }
1379    }
1380
1381    /// Wakes the guest for the specified channel if it's open and the rings are in a state where
1382    /// notification is expected.
1383    fn unstick_channel_by_id(&mut self, id: OfferInstanceId) {
1384        let Some(version) = self.server.get_version() else {
1385            tracelimit::warn_ratelimited!("cannot unstick when not connected");
1386            return;
1387        };
1388
1389        if let Some(channel) = self.inner.channels.get_mut(&id.offer_id) {
1390            if channel.seq != id.seq {
1391                // The channel was revoked.
1392                return;
1393            }
1394
1395            // The channel was closed and reopened before the delay expired, so wait again to ensure
1396            // we don't signal too early.
1397            if channel.unstick_state == ChannelUnstickState::NeedsRequeue {
1398                channel.unstick_state = ChannelUnstickState::Queued;
1399                let mut timer = PolledTimer::new(&self.driver);
1400                let delay = self.channel_unstick_delay.unwrap();
1401                self.channel_unstickers.push(Box::pin(async move {
1402                    timer.sleep(delay).await;
1403                    id
1404                }));
1405
1406                return;
1407            }
1408
1409            channel.unstick_state = ChannelUnstickState::None;
1410            let gm = select_gm_for_channel(
1411                &self.inner.gm,
1412                self.inner.private_gm.as_ref(),
1413                version,
1414                channel,
1415            );
1416            if let Err(err) = Self::unstick_channel(gm, channel, false, false) {
1417                tracelimit::warn_ratelimited!(
1418                    channel = %channel.key,
1419                    error = err.as_ref() as &dyn std::error::Error,
1420                    "could not unstick channel"
1421                );
1422            }
1423        }
1424    }
1425
1426    fn unstick_channel(
1427        gm: &GuestMemory,
1428        channel: &Channel,
1429        force: bool,
1430        unstick_host: bool,
1431    ) -> anyhow::Result<()> {
1432        if let ChannelState::Open(state) = &channel.state {
1433            if force {
1434                tracing::info!(channel = %channel.key, "waking host and guest");
1435                if unstick_host {
1436                    channel.guest_to_host_event.0.deliver();
1437                }
1438                state.host_to_guest_interrupt.deliver();
1439                return Ok(());
1440            }
1441
1442            let gpadl = channel
1443                .gpadls
1444                .clone()
1445                .view()
1446                .map(state.open_params.open_data.ring_gpadl_id)
1447                .context("couldn't find ring gpadl")?;
1448
1449            let aligned = AlignedGpadlView::new(gpadl)
1450                .ok()
1451                .context("ring not aligned")?;
1452            let (in_gpadl, out_gpadl) = aligned
1453                .split(state.open_params.open_data.ring_offset)
1454                .ok()
1455                .context("couldn't split ring")?;
1456
1457            if let Err(err) = Self::unstick_incoming_ring(
1458                gm,
1459                channel,
1460                in_gpadl,
1461                unstick_host.then_some(channel.guest_to_host_event.as_ref()),
1462                &state.host_to_guest_interrupt,
1463            ) {
1464                tracelimit::warn_ratelimited!(
1465                    channel = %channel.key,
1466                    error = err.as_ref() as &dyn std::error::Error,
1467                    "could not unstick incoming ring"
1468                );
1469            }
1470            if let Err(err) = Self::unstick_outgoing_ring(
1471                gm,
1472                channel,
1473                out_gpadl,
1474                unstick_host.then_some(channel.guest_to_host_event.as_ref()),
1475                &state.host_to_guest_interrupt,
1476            ) {
1477                tracelimit::warn_ratelimited!(
1478                    channel = %channel.key,
1479                    error = err.as_ref() as &dyn std::error::Error,
1480                    "could not unstick outgoing ring"
1481                );
1482            }
1483        }
1484        Ok(())
1485    }
1486
1487    fn unstick_incoming_ring(
1488        gm: &GuestMemory,
1489        channel: &Channel,
1490        in_gpadl: AlignedGpadlView,
1491        guest_to_host_event: Option<&ChannelEvent>,
1492        host_to_guest_interrupt: &Interrupt,
1493    ) -> anyhow::Result<()> {
1494        let control_page = lock_gpn_with_subrange(gm, in_gpadl.gpns()[0])?;
1495        if let Some(guest_to_host_event) = guest_to_host_event {
1496            if ring::reader_needs_signal(control_page.pages()[0]) {
1497                tracelimit::info_ratelimited!(channel = %channel.key, "waking host for incoming ring");
1498                guest_to_host_event.0.deliver();
1499            }
1500        }
1501
1502        let ring_size = gpadl_ring_size(&in_gpadl).try_into()?;
1503        if ring::writer_needs_signal(control_page.pages()[0], ring_size) {
1504            tracelimit::info_ratelimited!(channel = %channel.key, "waking guest for incoming ring");
1505            host_to_guest_interrupt.deliver();
1506        }
1507        Ok(())
1508    }
1509
1510    fn unstick_outgoing_ring(
1511        gm: &GuestMemory,
1512        channel: &Channel,
1513        out_gpadl: AlignedGpadlView,
1514        guest_to_host_event: Option<&ChannelEvent>,
1515        host_to_guest_interrupt: &Interrupt,
1516    ) -> anyhow::Result<()> {
1517        let control_page = lock_gpn_with_subrange(gm, out_gpadl.gpns()[0])?;
1518        if ring::reader_needs_signal(control_page.pages()[0]) {
1519            tracelimit::info_ratelimited!(channel = %channel.key, "waking guest for outgoing ring");
1520            host_to_guest_interrupt.deliver();
1521        }
1522
1523        if let Some(guest_to_host_event) = guest_to_host_event {
1524            let ring_size = gpadl_ring_size(&out_gpadl).try_into()?;
1525            if ring::writer_needs_signal(control_page.pages()[0], ring_size) {
1526                tracelimit::info_ratelimited!(channel = %channel.key, "waking host for outgoing ring");
1527                guest_to_host_event.0.deliver();
1528            }
1529        }
1530        Ok(())
1531    }
1532}
1533
1534impl Notifier for ServerTaskInner {
1535    fn notify(&mut self, offer_id: OfferId, action: channels::Action) {
1536        let channel = self
1537            .channels
1538            .get_mut(&offer_id)
1539            .expect("channel does not exist");
1540
1541        fn handle<I: 'static + Send, R: 'static + Send>(
1542            offer_id: OfferId,
1543            channel: &Channel,
1544            req: impl FnOnce(Rpc<I, R>) -> ChannelRequest,
1545            input: I,
1546            f: impl 'static + Send + FnOnce(R) -> ChannelResponse,
1547        ) -> Pin<Box<dyn Send + Future<Output = (OfferId, u64, Result<ChannelResponse, RpcError>)>>>
1548        {
1549            let recv = channel.send.call(req, input);
1550            let seq = channel.seq;
1551            Box::pin(async move {
1552                let r = recv.await.map(f);
1553                (offer_id, seq, r)
1554            })
1555        }
1556
1557        let response = match action {
1558            channels::Action::Open(open_params, version) => {
1559                let seq = channel.seq;
1560                let key = channel.key;
1561                match self.open_channel(offer_id, &open_params) {
1562                    Ok((channel, interrupt)) => handle(
1563                        offer_id,
1564                        channel,
1565                        ChannelRequest::Open,
1566                        OpenRequest::new(
1567                            open_params.open_data,
1568                            interrupt,
1569                            version.feature_flags,
1570                            channel.flags,
1571                        ),
1572                        ChannelResponse::Open,
1573                    ),
1574                    Err(err) => {
1575                        tracelimit::error_ratelimited!(
1576                            err = err.as_ref() as &dyn std::error::Error,
1577                            ?offer_id,
1578                            %key,
1579                            "could not open channel",
1580                        );
1581
1582                        // Return an error response to the channels module if the open_channel call
1583                        // failed.
1584                        Box::pin(future::ready((
1585                            offer_id,
1586                            seq,
1587                            Ok(ChannelResponse::Open(false)),
1588                        )))
1589                    }
1590                }
1591            }
1592            channels::Action::Close => {
1593                if let Some(channel_bitmap) = self.channel_bitmap.as_ref() {
1594                    if let ChannelState::Open(ref state) = channel.state {
1595                        channel_bitmap.unregister_channel(state.open_params.event_flag);
1596                    }
1597                }
1598
1599                channel.state = ChannelState::Closing;
1600                handle(offer_id, channel, ChannelRequest::Close, (), |()| {
1601                    ChannelResponse::Close
1602                })
1603            }
1604            channels::Action::Gpadl(gpadl_id, count, buf) => {
1605                channel.gpadls.add(
1606                    gpadl_id,
1607                    MultiPagedRangeBuf::from_range_buffer(count.into(), buf.clone()).unwrap(),
1608                );
1609                handle(
1610                    offer_id,
1611                    channel,
1612                    ChannelRequest::Gpadl,
1613                    GpadlRequest {
1614                        id: gpadl_id,
1615                        count,
1616                        buf,
1617                    },
1618                    move |r| ChannelResponse::Gpadl(gpadl_id, r),
1619                )
1620            }
1621            channels::Action::TeardownGpadl {
1622                gpadl_id,
1623                post_restore,
1624            } => {
1625                if !post_restore {
1626                    channel.gpadls.remove(gpadl_id, Box::new(|| ()));
1627                }
1628
1629                handle(
1630                    offer_id,
1631                    channel,
1632                    ChannelRequest::TeardownGpadl,
1633                    gpadl_id,
1634                    move |()| ChannelResponse::TeardownGpadl(gpadl_id),
1635                )
1636            }
1637            channels::Action::Modify { target_vp } => {
1638                let ChannelState::Open(state) = &mut channel.state else {
1639                    unreachable!();
1640                };
1641
1642                if let Err(err) = state.set_event_port_target_vp(target_vp) {
1643                    tracelimit::error_ratelimited!(
1644                        error = err.as_ref() as &dyn std::error::Error,
1645                        channel = %channel.key,
1646                        "could not modify channel",
1647                    );
1648
1649                    // Send an immediate error response.
1650                    let seq = channel.seq;
1651                    Box::pin(async move {
1652                        (
1653                            offer_id,
1654                            seq,
1655                            Ok(ChannelResponse::Modify(protocol::STATUS_UNSUCCESSFUL)),
1656                        )
1657                    })
1658                } else {
1659                    handle(
1660                        offer_id,
1661                        channel,
1662                        ChannelRequest::Modify,
1663                        ModifyRequest::TargetVp { target_vp },
1664                        ChannelResponse::Modify,
1665                    )
1666                }
1667            }
1668        };
1669        self.channel_responses.push(response);
1670    }
1671
1672    fn modify_connection(&mut self, mut request: ModifyConnectionRequest) -> anyhow::Result<()> {
1673        self.map_interrupt_page(request.interrupt_page)
1674            .context("Failed to map interrupt page.")?;
1675
1676        self.set_monitor_page(&mut request)
1677            .context("Failed to map monitor page.")?;
1678
1679        if let Some(vp) = request.target_message_vp {
1680            self.message_port.set_target_vp(vp)?;
1681        }
1682
1683        if request.notify_relay {
1684            self.relay_send.send(request.into());
1685        }
1686
1687        Ok(())
1688    }
1689
1690    fn forward_unhandled(&mut self, request: InitiateContactRequest) {
1691        if let Some(external_server) = &self.external_server_send {
1692            external_server.send(request);
1693        } else {
1694            tracing::warn!(?request, "nowhere to forward unhandled request")
1695        }
1696    }
1697
1698    fn inspect(&self, version: Option<VersionInfo>, offer_id: OfferId, req: inspect::Request<'_>) {
1699        let channel = self.channels.get(&offer_id).expect("should exist");
1700        let mut resp = req.respond();
1701        // Only inspect ring buffers if we have an active connection; during
1702        // disconnect/reset, `version` may be None even if an individual channel
1703        // is still in the Open state, and we must not panic during inspect
1704        // (e.g., timeout diagnostics rely on inspect succeeding).
1705        if let (ChannelState::Open(state), Some(version)) = (&channel.state, version) {
1706            let mem = self.get_gm_for_channel(version, channel);
1707            inspect_rings(
1708                &mut resp,
1709                mem,
1710                channel.gpadls.clone(),
1711                &state.open_params.open_data,
1712            );
1713        }
1714    }
1715
1716    fn send_message(&mut self, message: &OutgoingMessage, target: MessageTarget) -> bool {
1717        // If the server is paused, queue all messages, to avoid affecting synic
1718        // state during/after it has been saved or reset.
1719        //
1720        // Note that messages to reserved channels or custom targets will be
1721        // dropped. However, such messages should only be sent in response to
1722        // guest requests, which should not be processed while the server is
1723        // paused.
1724        //
1725        // FUTURE: it would be better to ensure that no messages are generated
1726        // by operations that run while the server is paused. E.g., defer
1727        // sending offer or revoke messages for new or revoked offers. This
1728        // would prevent the queue from growing without bound.
1729        if !self.running && !self.send_messages_while_stopped {
1730            if !matches!(target, MessageTarget::Default) {
1731                tracelimit::error_ratelimited!(?target, "dropping message while paused");
1732            }
1733            return false;
1734        }
1735
1736        let mut port_storage;
1737        let port = match target {
1738            MessageTarget::Default => self.message_port.as_mut(),
1739            MessageTarget::ReservedChannel(offer_id, target) => {
1740                if let Some(port) = self.get_reserved_channel_message_port(offer_id, target) {
1741                    port.as_mut()
1742                } else {
1743                    // Updating the port failed, so there is no way to send the message.
1744                    return true;
1745                }
1746            }
1747            MessageTarget::Custom(target) => {
1748                port_storage = match self.synic.new_guest_message_port(
1749                    self.redirect_vtl,
1750                    target.vp,
1751                    target.sint,
1752                ) {
1753                    Ok(port) => port,
1754                    Err(err) => {
1755                        tracing::error!(
1756                            ?err,
1757                            ?self.redirect_vtl,
1758                            ?target,
1759                            "could not create message port"
1760                        );
1761
1762                        // There is no way to send the message.
1763                        return true;
1764                    }
1765                };
1766                port_storage.as_mut()
1767            }
1768        };
1769
1770        // If this returns Pending, the channels module will queue the message and the ServerTask
1771        // main loop will try to send it again later.
1772        matches!(
1773            port.poll_post_message(
1774                &mut std::task::Context::from_waker(std::task::Waker::noop()),
1775                VMBUS_MESSAGE_TYPE,
1776                message.data()
1777            ),
1778            Poll::Ready(())
1779        )
1780    }
1781
1782    fn notify_hvsock(&mut self, request: &HvsockConnectRequest) {
1783        tracing::debug!(?request, "received hvsock connect request");
1784        self.hvsock_requests += 1;
1785        self.hvsock_send.send(*request);
1786    }
1787
1788    fn reset_complete(&mut self) {
1789        if let Some(monitor) = self.synic.monitor_support() {
1790            if let Err(err) = monitor.set_monitor_page(self.vtl, None) {
1791                tracing::warn!(?err, "resetting monitor page failed")
1792            }
1793        }
1794
1795        self.unreserve_channels();
1796        for done in self.reset_done.drain(..) {
1797            done.complete(());
1798        }
1799    }
1800
1801    fn unload_complete(&mut self) {
1802        self.unreserve_channels();
1803    }
1804}
1805
1806impl ServerTaskInner {
1807    fn open_channel(
1808        &mut self,
1809        offer_id: OfferId,
1810        open_params: &OpenParams,
1811    ) -> anyhow::Result<(&mut Channel, Interrupt)> {
1812        let channel = self
1813            .channels
1814            .get_mut(&offer_id)
1815            .expect("channel does not exist");
1816
1817        // Always register with the channel bitmap; if Win7, this may be unnecessary.
1818        if let Some(channel_bitmap) = self.channel_bitmap.as_ref() {
1819            channel_bitmap.register_channel(
1820                open_params.event_flag,
1821                channel.guest_to_host_event.0.clone(),
1822            );
1823        }
1824        // Always set up an event port; if V1, this will be unused.
1825        // N.B. The event port must be created before the device is notified of the open by the
1826        //      caller. The device may begin communicating with the guest immediately when it is
1827        //      notified, so the event port must exist so that the guest can send interrupts.
1828        let event_port = self
1829            .synic
1830            .add_event_port(
1831                open_params.connection_id,
1832                self.vtl,
1833                channel.guest_to_host_event.clone(),
1834                open_params.monitor_info,
1835            )
1836            .context("failed to create guest-to-host event port")?;
1837
1838        // For pre-Win8 guests, the host-to-guest event always targets vp 0 and the channel
1839        // bitmap is used instead of the event flag.
1840        let (target_vp, event_flag) = if self.channel_bitmap.is_some() {
1841            (Some(0), 0)
1842        } else {
1843            (open_params.open_data.target_vp, open_params.event_flag)
1844        };
1845
1846        let (guest_event_port, interrupt) = if let Some(target_vp) = target_vp {
1847            let (target_vtl, target_sint) = if open_params.flags.redirect_interrupt() {
1848                (self.redirect_vtl, self.redirect_sint)
1849            } else {
1850                (self.vtl, VMBUS_SINT)
1851            };
1852
1853            let guest_event_port = self.synic.new_guest_event_port(
1854                VmbusServer::get_child_event_port_id(open_params.channel_id, VMBUS_SINT, self.vtl),
1855                target_vtl,
1856                target_vp,
1857                target_sint,
1858                event_flag,
1859                open_params.monitor_info,
1860            )?;
1861
1862            let interrupt = ChannelBitmap::create_interrupt(
1863                &self.channel_bitmap,
1864                guest_event_port.interrupt(),
1865                open_params.event_flag,
1866            );
1867
1868            (Some(guest_event_port), interrupt)
1869        } else {
1870            // Use a dummy interrupt which does nothing.
1871            (None, Interrupt::null())
1872        };
1873
1874        // Delete any previously reserved state.
1875        channel.reserved_state.message_port = None;
1876
1877        // If the channel is reserved, create a message port for it.
1878        if let Some(target) = open_params.reserved_target {
1879            channel.reserved_state.message_port = Some(self.synic.new_guest_message_port(
1880                self.redirect_vtl,
1881                target.vp,
1882                target.sint,
1883            )?);
1884
1885            channel.reserved_state.target = target;
1886        }
1887
1888        channel.state = ChannelState::Open(Box::new(ChannelOpenState {
1889            open_params: *open_params,
1890            _event_port: event_port,
1891            guest_event_port,
1892            host_to_guest_interrupt: interrupt.clone(),
1893        }));
1894        Ok((channel, interrupt))
1895    }
1896
1897    /// If the client specified an interrupt page, map it into host memory and
1898    /// set up the shared event port.
1899    fn map_interrupt_page(&mut self, interrupt_page: Update<u64>) -> anyhow::Result<()> {
1900        let interrupt_page = match interrupt_page {
1901            Update::Unchanged => return Ok(()),
1902            Update::Reset => {
1903                self.channel_bitmap = None;
1904                self.shared_event_port = None;
1905                return Ok(());
1906            }
1907            Update::Set(interrupt_page) => interrupt_page,
1908        };
1909
1910        assert_ne!(interrupt_page, 0);
1911
1912        if interrupt_page % PAGE_SIZE as u64 != 0 {
1913            anyhow::bail!("interrupt page {:#x} is not page aligned", interrupt_page);
1914        }
1915
1916        // Use a subrange to access the interrupt page to give GuestMemory's without a full mapping
1917        // a chance to create one.
1918        let interrupt_page = lock_page_with_subrange(&self.gm, interrupt_page)?;
1919        let channel_bitmap = Arc::new(ChannelBitmap::new(interrupt_page));
1920        self.channel_bitmap = Some(channel_bitmap.clone());
1921
1922        // Create the shared event port for pre-Win8 guests.
1923        let interrupt = Interrupt::from_fn(move || {
1924            channel_bitmap.handle_shared_interrupt();
1925        });
1926
1927        self.shared_event_port = Some(self.synic.add_event_port(
1928            SHARED_EVENT_CONNECTION_ID,
1929            self.vtl,
1930            Arc::new(ChannelEvent(interrupt)),
1931            None,
1932        )?);
1933
1934        Ok(())
1935    }
1936
1937    fn set_monitor_page(&mut self, request: &mut ModifyConnectionRequest) -> anyhow::Result<()> {
1938        let monitor_page = match request.monitor_page {
1939            Update::Unchanged => return Ok(()),
1940            Update::Reset => None,
1941            Update::Set(value) => Some(value),
1942        };
1943
1944        // TODO: can this check be moved into channels.rs?
1945        if self.channels.iter().any(|(_, c)| {
1946            matches!(
1947                &c.state,
1948                ChannelState::Open(state) if state.open_params.monitor_info.is_some()
1949            )
1950        }) {
1951            anyhow::bail!("attempt to change monitor page while open channels using mnf");
1952        }
1953
1954        // Check if the server is handling MNF.
1955        // N.B. If the server is not handling MNF, there is currently no way to request
1956        //      server-allocated monitor pages from the relay host.
1957        if let Some(mnf_support) = self.mnf_support.as_mut() {
1958            if let Some(monitor) = self.synic.monitor_support() {
1959                mnf_support.allocated_monitor_page = None;
1960
1961                if let Some(version) = request.version {
1962                    if version.feature_flags.server_specified_monitor_pages() {
1963                        if let Some(monitor_page) = monitor.allocate_monitor_page(self.vtl)? {
1964                            tracelimit::info_ratelimited!(
1965                                ?monitor_page,
1966                                "using server-allocated monitor pages"
1967                            );
1968                            mnf_support.allocated_monitor_page = Some(monitor_page);
1969                        }
1970                    }
1971                }
1972
1973                // If no monitor page was allocated above, use the one provided by the client.
1974                if mnf_support.allocated_monitor_page.is_none() {
1975                    if let Err(err) = monitor.set_monitor_page(self.vtl, monitor_page) {
1976                        anyhow::bail!(
1977                            "setting monitor page failed, err = {err:?}, monitor_page = {monitor_page:?}"
1978                        );
1979                    }
1980                }
1981            }
1982
1983            // If MNF is configured to be handled by this server (even if it's not actually
1984            // supported by the synic), don't forward the pages to the relay.
1985            request.monitor_page = Update::Unchanged;
1986        }
1987
1988        Ok(())
1989    }
1990
1991    fn get_reserved_channel_message_port(
1992        &mut self,
1993        offer_id: OfferId,
1994        new_target: ConnectionTarget,
1995    ) -> Option<&mut Box<dyn GuestMessagePort>> {
1996        let channel = self
1997            .channels
1998            .get_mut(&offer_id)
1999            .expect("channel does not exist");
2000
2001        assert!(
2002            channel.reserved_state.message_port.is_some(),
2003            "channel is not reserved"
2004        );
2005
2006        // On close, the guest may have changed the message target it wants to use for the close
2007        // response. If so, update the message port.
2008        if channel.reserved_state.target.sint != new_target.sint {
2009            // Destroy the old port before creating the new one.
2010            channel.reserved_state.message_port = None;
2011            let message_port = self
2012                .synic
2013                .new_guest_message_port(self.redirect_vtl, new_target.vp, new_target.sint)
2014                .inspect_err(|err| {
2015                    tracing::error!(
2016                        key = %channel.key,
2017                        ?err,
2018                        ?self.redirect_vtl,
2019                        ?new_target,
2020                        "could not create reserved channel message port"
2021                    )
2022                })
2023                .ok()?;
2024
2025            channel.reserved_state.message_port = Some(message_port);
2026            channel.reserved_state.target = new_target;
2027        } else if channel.reserved_state.target.vp != new_target.vp {
2028            let message_port = channel.reserved_state.message_port.as_mut().unwrap();
2029
2030            // The vp has changed, but the SINT is the same. Just update the vp. If this fails,
2031            // ignore it and just send to the old vp.
2032            if let Err(err) = message_port.set_target_vp(new_target.vp) {
2033                tracing::error!(
2034                    key = %channel.key,
2035                    ?err,
2036                    ?self.redirect_vtl,
2037                    ?new_target,
2038                    "could not update reserved channel message port"
2039                );
2040            }
2041
2042            channel.reserved_state.target = new_target;
2043            return Some(message_port);
2044        }
2045
2046        Some(channel.reserved_state.message_port.as_mut().unwrap())
2047    }
2048
2049    fn unreserve_channels(&mut self) {
2050        // Unreserve all closed channels.
2051        for channel in self.channels.values_mut() {
2052            if let ChannelState::Closed = channel.state {
2053                channel.reserved_state.message_port = None;
2054            }
2055        }
2056    }
2057
2058    fn get_gm_for_channel(&self, version: VersionInfo, channel: &Channel) -> &GuestMemory {
2059        select_gm_for_channel(&self.gm, self.private_gm.as_ref(), version, channel)
2060    }
2061}
2062
2063fn select_gm_for_channel<'a>(
2064    gm: &'a GuestMemory,
2065    private_gm: Option<&'a GuestMemory>,
2066    version: VersionInfo,
2067    channel: &Channel,
2068) -> &'a GuestMemory {
2069    if channel.flags.confidential_ring_buffer() && version.feature_flags.confidential_channels() {
2070        if let Some(private_gm) = private_gm {
2071            return private_gm;
2072        }
2073    }
2074
2075    gm
2076}
2077
2078/// Control point for [`VmbusServer`], allowing callers to offer channels.
2079#[derive(Clone)]
2080pub struct VmbusServerControl {
2081    mem: GuestMemory,
2082    private_mem: Option<GuestMemory>,
2083    send: mesh::Sender<OfferRequest>,
2084    use_event: bool,
2085    force_confidential_external_memory: bool,
2086    force_gpa_pinning: bool,
2087}
2088
2089impl VmbusServerControl {
2090    /// Offers a channel to the vmbus server, where the flags and user_defined data are already set.
2091    /// This is used by the relay to forward the host's parameters.
2092    pub async fn offer_core(&self, mut offer_info: OfferInfo) -> anyhow::Result<OfferResources> {
2093        if self.force_gpa_pinning {
2094            tracing::warn!(
2095                key = %offer_info.params.key(),
2096                "forcing GPA pinning for channel"
2097            );
2098
2099            offer_info
2100                .params
2101                .flags
2102                .set_require_pinned_external_memory(true);
2103        }
2104
2105        let flags = offer_info.params.flags;
2106        self.send
2107            .call_failable(OfferRequest::Offer, offer_info)
2108            .await?;
2109        Ok(OfferResources::new(
2110            self.mem.clone(),
2111            if flags.confidential_ring_buffer() || flags.confidential_external_memory() {
2112                self.private_mem.clone()
2113            } else {
2114                None
2115            },
2116        ))
2117    }
2118
2119    /// Force reset all channels and protocol state, without requiring the
2120    /// server to be paused.
2121    pub async fn force_reset(&self) -> anyhow::Result<()> {
2122        self.send
2123            .call(OfferRequest::ForceReset, ())
2124            .await
2125            .context("vmbus server is gone")
2126    }
2127
2128    async fn offer(&self, request: OfferInput) -> anyhow::Result<OfferResources> {
2129        let mut offer_info = OfferInfo {
2130            params: request.params.into(),
2131            event: request.event,
2132            request_send: request.request_send,
2133            server_request_recv: request.server_request_recv,
2134        };
2135
2136        if self.force_confidential_external_memory {
2137            tracing::warn!(
2138                key = %offer_info.params.key(),
2139                "forcing confidential external memory for channel"
2140            );
2141
2142            offer_info
2143                .params
2144                .flags
2145                .set_confidential_external_memory(true);
2146        }
2147
2148        self.offer_core(offer_info).await
2149    }
2150}
2151
2152/// Inspects the specified ring buffer state by directly accessing guest memory.
2153fn inspect_rings(
2154    resp: &mut inspect::Response<'_>,
2155    gm: &GuestMemory,
2156    gpadl_map: Arc<GpadlMap>,
2157    open_data: &OpenData,
2158) -> Option<()> {
2159    let gpadl = gpadl_map
2160        .view()
2161        .map(GpadlId(open_data.ring_gpadl_id.0))
2162        .ok()?;
2163
2164    let aligned = AlignedGpadlView::new(gpadl).ok()?;
2165    let (in_gpadl, out_gpadl) = aligned.split(open_data.ring_offset).ok()?;
2166    resp.child("incoming_ring", |req| inspect_ring(req, &in_gpadl, gm));
2167    resp.child("outgoing_ring", |req| inspect_ring(req, &out_gpadl, gm));
2168    Some(())
2169}
2170
2171/// Inspects the incoming or outgoing ring buffer by directly accessing guest memory.
2172fn inspect_ring(req: inspect::Request<'_>, gpadl: &AlignedGpadlView, gm: &GuestMemory) {
2173    let mut resp = req.respond();
2174
2175    resp.hex("ring_size", gpadl_ring_size(gpadl));
2176
2177    // Lock just the control page. Use a subrange to allow a GuestMemory without a full mapping to
2178    // create one.
2179    if let Ok(pages) = lock_gpn_with_subrange(gm, gpadl.gpns()[0]) {
2180        ring::inspect_ring(pages.pages()[0], &mut resp);
2181    }
2182}
2183
2184fn gpadl_ring_size(gpadl: &AlignedGpadlView) -> usize {
2185    // Data size excluding the control page.
2186    (gpadl.gpns().len() - 1) * PAGE_SIZE
2187}
2188
2189/// Helper to create a subrange before locking a single page.
2190///
2191/// This allows us to lock a page in a `GuestMemory` that doesn't have a full mapping, but can
2192/// create one for a subrange.
2193fn lock_page_with_subrange(gm: &GuestMemory, offset: u64) -> anyhow::Result<guestmem::LockedPages> {
2194    Ok(gm.lockable_subrange(offset, PAGE_SIZE as u64)?.lock_gpns(
2195        guestmem::AccessType::Write,
2196        false,
2197        &[0],
2198    )?)
2199}
2200
2201/// Helper to create a subrange before locking a single page from a gpn.
2202///
2203/// This allows us to lock a page in a `GuestMemory` that doesn't have a full mapping, but can
2204/// create one for a subrange.
2205fn lock_gpn_with_subrange(gm: &GuestMemory, gpn: u64) -> anyhow::Result<guestmem::LockedPages> {
2206    lock_page_with_subrange(gm, gpn * PAGE_SIZE as u64)
2207}
2208
2209pub(crate) struct MessageSender {
2210    send: mpsc::Sender<SynicMessage>,
2211    multiclient: bool,
2212}
2213
2214impl MessageSender {
2215    fn poll_handle_message(
2216        &self,
2217        cx: &mut std::task::Context<'_>,
2218        msg: &[u8],
2219        trusted: bool,
2220    ) -> Poll<Result<(), SendError>> {
2221        let mut send = self.send.clone();
2222        ready!(send.poll_ready(cx))?;
2223        send.start_send(SynicMessage {
2224            data: msg.to_vec(),
2225            multiclient: self.multiclient,
2226            trusted,
2227        })?;
2228
2229        Poll::Ready(Ok(()))
2230    }
2231}
2232
2233impl MessagePort for MessageSender {
2234    fn poll_handle_message(
2235        &self,
2236        cx: &mut std::task::Context<'_>,
2237        msg: &[u8],
2238        trusted: bool,
2239    ) -> Poll<()> {
2240        if let Err(err) = ready!(self.poll_handle_message(cx, msg, trusted)) {
2241            tracelimit::error_ratelimited!(
2242                error = &err as &dyn std::error::Error,
2243                "failed to send message"
2244            );
2245        }
2246
2247        Poll::Ready(())
2248    }
2249}
2250
2251#[async_trait]
2252impl ParentBus for VmbusServerControl {
2253    async fn add_child(&self, request: OfferInput) -> anyhow::Result<OfferResources> {
2254        self.offer(request).await
2255    }
2256
2257    fn clone_bus(&self) -> Box<dyn ParentBus> {
2258        Box::new(self.clone())
2259    }
2260
2261    fn use_event(&self) -> bool {
2262        self.use_event
2263    }
2264}