1#![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
15pub 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)]
150pub enum SavedStateRequest {
152 Set(FailableRpc<Box<channels::SavedState>, ()>),
153 Clear(Rpc<(), ()>),
154}
155
156pub struct ServerChannelHalf<Request, Response> {
158 request_send: mesh::Sender<Request>,
159 response_receive: mesh::Receiver<Response>,
160}
161
162pub struct RelayChannelHalf<Request, Response> {
164 pub request_receive: mesh::Receiver<Request>,
165 pub response_send: mesh::Sender<Response>,
166}
167
168pub 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 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#[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#[derive(Debug, Copy, Clone)]
216pub enum ModifyRelayResponse {
217 Supported(protocol::ConnectionState, protocol::FeatureFlags),
221 Unsupported,
224 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 #[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 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 pub fn private_gm(mut self, private_gm: Option<GuestMemory>) -> Self {
328 self.private_gm = private_gm;
329 self
330 }
331
332 pub fn vtl(mut self, vtl: Vtl) -> Self {
334 self.vtl = vtl;
335 self
336 }
337
338 pub fn hvsock_notify(mut self, hvsock_notify: Option<HvsockServerChannelHalf>) -> Self {
340 self.hvsock_notify = hvsock_notify;
341 self
342 }
343
344 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 pub fn server_relay(mut self, server_relay: Option<VmbusServerChannelHalf>) -> Self {
356 self.server_relay = server_relay;
357 self
358 }
359
360 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 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 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 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 pub fn max_version(mut self, max_version: Option<MaxVersionInfo>) -> Self {
398 self.max_version = max_version;
399 self
400 }
401
402 pub fn delay_max_version(mut self, delay: bool) -> Self {
407 self.delay_max_version = delay;
408 self
409 }
410
411 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 pub fn enable_mnf(mut self, enable: bool) -> Self {
425 self.enable_mnf = enable;
426 self
427 }
428
429 pub fn force_confidential_external_memory(mut self, force: bool) -> Self {
432 self.force_confidential_external_memory = force;
433 self
434 }
435
436 pub fn support_gpa_pinning(mut self, support: bool) -> Self {
439 self.support_gpa_pinning = support;
440 self
441 }
442
443 pub fn force_gpa_pinning(mut self, force: bool) -> Self {
445 self.force_gpa_pinning = force;
446 self
447 }
448
449 pub fn send_messages_while_stopped(mut self, send: bool) -> Self {
456 self.send_messages_while_stopped = send;
457 self
458 }
459
460 pub fn channel_unstick_delay(mut self, delay: Option<Duration>) -> Self {
467 self.channel_unstick_delay = delay;
468 self
469 }
470
471 pub fn use_absolute_channel_order(mut self, assign: bool) -> Self {
475 self.use_absolute_channel_order = assign;
476 self
477 }
478
479 pub fn build(self) -> anyhow::Result<VmbusServer> {
484 #[expect(clippy::disallowed_methods)] 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 let connection_id = if self.vtl == Vtl::Vtl0 && !self.use_message_redirect {
500 MESSAGE_CONNECTION_ID
501 } else {
502 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 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 server.set_require_server_allocated_mnf(self.enable_mnf && self.private_gm.is_some());
557
558 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 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 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 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 pub async fn stop(&self) {
683 self.task_send.call(VmbusRequest::Stop, ()).await.unwrap()
684 }
685
686 pub fn start(&self) {
688 self.task_send.send(VmbusRequest::Start);
689 }
690
691 pub async fn reset(&self) {
693 tracing::debug!("resetting channel state");
694 self.task_send.call(VmbusRequest::Reset, ()).await.unwrap()
695 }
696
697 pub async fn shutdown(self) {
699 drop(self.task_send);
700 let _ = self.task.await;
701 }
702
703 pub fn control(&self) -> Arc<VmbusServerControl> {
705 self.control.clone()
706 }
707
708 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#[derive(Default)]
731struct MnfSupport {
732 allocated_monitor_page: Option<MonitorPageGpas>,
733}
734
735#[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_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 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 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 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 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 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 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 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 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, ¶ms)?;
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 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 let response = match response {
1186 ModifyRelayResponse::Supported(state, features) => {
1187 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 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 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 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 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 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 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! { 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 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 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 return;
1393 }
1394
1395 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 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 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 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 !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 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 return true;
1764 }
1765 };
1766 port_storage.as_mut()
1767 }
1768 };
1769
1770 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 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 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 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 (None, Interrupt::null())
1872 };
1873
1874 channel.reserved_state.message_port = None;
1876
1877 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 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 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 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 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 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 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 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 if channel.reserved_state.target.sint != new_target.sint {
2009 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 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 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#[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 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 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
2152fn 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
2171fn 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 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 (gpadl.gpns().len() - 1) * PAGE_SIZE
2187}
2188
2189fn 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
2201fn 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}