1#![expect(missing_docs)]
7#![forbid(unsafe_code)]
8
9pub mod driver;
10pub mod filter;
11mod hvsock;
12pub mod saved_state;
13
14pub use self::saved_state::SavedState;
15use anyhow::Context as _;
16use anyhow::Result;
17use futures::FutureExt;
18use futures::StreamExt;
19use futures::future::OptionFuture;
20use futures::stream::SelectAll;
21use futures_concurrency::future::Race;
22use guid::Guid;
23use inspect::Inspect;
24use mesh::rpc::FailableRpc;
25use mesh::rpc::Rpc;
26use mesh::rpc::RpcSend;
27use pal_async::task::Spawn;
28use pal_async::task::Task;
29use pal_event::Event;
30use std::collections::HashMap;
31use std::collections::VecDeque;
32use std::collections::hash_map;
33use std::convert::TryInto;
34use std::future::Future;
35use std::future::poll_fn;
36use std::ops::Deref;
37use std::ops::DerefMut;
38use std::pin::pin;
39use std::sync::Arc;
40use std::sync::atomic::AtomicU32;
41use std::sync::atomic::Ordering;
42use std::task::Context;
43use std::task::Poll;
44use thiserror::Error;
45use vmbus_async::async_dgram::AsyncRecv;
46use vmbus_async::async_dgram::AsyncRecvExt;
47use vmbus_channel::bus::GpadlRequest;
48use vmbus_channel::bus::ModifyRequest;
49use vmbus_channel::bus::OfferKey;
50use vmbus_channel::bus::OpenData;
51use vmbus_channel::gpadl::GpadlId;
52use vmbus_core::HvsockConnectRequest;
53use vmbus_core::OutgoingMessage;
54use vmbus_core::TaggedStream;
55use vmbus_core::VersionInfo;
56use vmbus_core::protocol;
57use vmbus_core::protocol::ChannelId;
58use vmbus_core::protocol::ConnectionState;
59use vmbus_core::protocol::FeatureFlags;
60use vmbus_core::protocol::Message;
61use vmbus_core::protocol::OpenChannelFlags;
62use vmbus_core::protocol::Version;
63use vmcore::interrupt::Interrupt;
64use vmcore::synic::MonitorPageGpas;
65use zerocopy::Immutable;
66use zerocopy::IntoBytes;
67use zerocopy::KnownLayout;
68
69const SINT: u8 = 2;
70const VTL: u8 = 0;
71const SUPPORTED_VERSIONS: &[Version] = &[Version::Iron, Version::Copper];
72const SUPPORTED_FEATURE_FLAGS: FeatureFlags = FeatureFlags::new()
73 .with_guest_specified_signal_parameters(true)
74 .with_channel_interrupt_redirection(true)
75 .with_modify_connection(true)
76 .with_client_id(true)
77 .with_pause_resume(true);
78
79pub trait SynicEventClient: Send + Sync {
81 fn map_event(&self, event_flag: u16, event: &Event) -> std::io::Result<()>;
83
84 fn unmap_event(&self, event_flag: u16);
86
87 fn signal_event(&self, connection_id: u32, event_flag: u16) -> std::io::Result<()>;
89}
90
91pub trait VmbusMessageSource: AsyncRecv + Send {
93 fn pause_message_stream(&mut self) {}
96
97 fn resume_message_stream(&mut self) {}
99}
100
101pub trait PollPostMessage: Send {
102 fn poll_post_message(
103 &mut self,
104 cx: &mut Context<'_>,
105 connection_id: u32,
106 typ: u32,
107 msg: &[u8],
108 ) -> Poll<()>;
109}
110
111#[derive(Inspect)]
112pub struct VmbusClient {
113 #[inspect(flatten, send = "TaskRequest::Inspect")]
114 task_send: mesh::Sender<TaskRequest>,
115 #[inspect(skip)]
116 access: VmbusClientAccess,
117 #[inspect(skip)]
118 task: Task<ClientTask>,
119}
120
121#[derive(Debug, thiserror::Error)]
122pub enum ConnectError {
123 #[error("invalid state to connect to the server")]
124 InvalidState,
125 #[error("no supported protocol versions")]
126 NoSupportedVersions,
127 #[error("failed to connect to the server: {0:?}")]
128 FailedToConnect(ConnectionState),
129}
130
131#[derive(Clone)]
132pub struct VmbusClientAccess {
133 client_request_send: mesh::Sender<ClientRequest>,
134}
135
136pub struct VmbusClientBuilder {
138 event_client: Arc<dyn SynicEventClient>,
139 msg_source: Box<dyn VmbusMessageSource>,
140 msg_client: Box<dyn PollPostMessage>,
141}
142
143impl VmbusClientBuilder {
144 pub fn new(
146 event_client: impl SynicEventClient + 'static,
147 msg_source: impl VmbusMessageSource + 'static,
148 msg_client: impl PollPostMessage + 'static,
149 ) -> Self {
150 Self {
151 event_client: Arc::new(event_client),
152 msg_source: Box::new(msg_source),
153 msg_client: Box::new(msg_client),
154 }
155 }
156
157 pub fn build(self, spawner: &impl Spawn) -> VmbusClient {
159 let (task_send, task_recv) = mesh::channel();
160 let (client_request_send, client_request_recv) = mesh::channel();
161
162 let inner = ClientTaskInner {
163 messages: OutgoingMessages {
164 poster: self.msg_client,
165 queued: VecDeque::new(),
166 state: OutgoingMessageState::Paused,
167 },
168 teardown_gpadls: HashMap::new(),
169 channel_requests: SelectAll::new(),
170 synic: SynicState {
171 event_flag_state: Vec::new(),
172 event_client: self.event_client,
173 },
174 };
175
176 let mut task = ClientTask {
177 inner,
178 channels: ChannelList::default(),
179 task_recv,
180 running: false,
181 msg_source: self.msg_source,
182 client_request_recv,
183 state: ClientState::Disconnected,
184 modify_request: None,
185 hvsock_tracker: hvsock::HvsockRequestTracker::new(),
186 };
187
188 let task = spawner.spawn("vmbus client", async move {
189 task.run().await;
190 task
191 });
192
193 VmbusClient {
194 access: VmbusClientAccess {
195 client_request_send,
196 },
197 task_send,
198 task,
199 }
200 }
201}
202
203impl VmbusClient {
204 pub async fn connect(
207 &mut self,
208 target_message_vp: u32,
209 monitor_page: Option<MonitorPageGpas>,
210 client_id: Guid,
211 ) -> Result<ConnectResult, ConnectError> {
212 let request = ConnectRequest {
213 target_message_vp,
214 monitor_page,
215 client_id,
216 };
217
218 self.access
219 .client_request_send
220 .call(ClientRequest::Connect, request)
221 .await
222 .unwrap()
223 }
224
225 pub async fn unload(self) {
226 self.access
227 .client_request_send
228 .call(ClientRequest::Unload, ())
229 .await
230 .unwrap();
231
232 self.sever().await;
233 }
234
235 pub fn access(&self) -> &VmbusClientAccess {
236 &self.access
237 }
238
239 pub fn start(&mut self) {
240 self.task_send.send(TaskRequest::Start);
241 }
242
243 pub async fn stop(&mut self) {
244 self.task_send
245 .call(TaskRequest::Stop, ())
246 .await
247 .expect("Failed to send stop request");
248 }
249
250 pub async fn save(&self) -> SavedState {
251 self.task_send
252 .call(TaskRequest::Save, ())
253 .await
254 .expect("Failed to send save request")
255 }
256
257 pub async fn restore(
258 &mut self,
259 state: SavedState,
260 ) -> Result<Option<ConnectResult>, RestoreError> {
261 self.task_send
262 .call(TaskRequest::Restore, state)
263 .await
264 .expect("Failed to send restore request")
265 }
266
267 pub async fn post_restore(&mut self) {
268 self.task_send
269 .call(TaskRequest::PostRestore, ())
270 .await
271 .expect("Failed to send post-restore request");
272 }
273
274 async fn sever(self) -> VmbusClientBuilder {
275 drop(self.task_send);
276 let task = self.task.await;
277 VmbusClientBuilder {
278 event_client: task.inner.synic.event_client,
279 msg_source: task.msg_source,
280 msg_client: task.inner.messages.poster,
281 }
282 }
283}
284
285#[derive(Debug)]
286pub struct ConnectResult {
287 pub version: VersionInfo,
288 pub offers: Vec<OfferInfo>,
289 pub offer_recv: mesh::Receiver<OfferInfo>,
290}
291
292impl VmbusClientAccess {
293 pub async fn modify(&self, request: ModifyConnectionRequest) -> ConnectionState {
294 self.client_request_send
295 .call(ClientRequest::Modify, request)
296 .await
297 .expect("Failed to send modify request")
298 }
299
300 pub fn connect_hvsock(
301 &self,
302 request: HvsockConnectRequest,
303 ) -> impl Future<Output = Option<OfferInfo>> + use<> {
304 self.client_request_send
305 .call(ClientRequest::HvsockConnect, request)
306 .map(|r| r.ok().flatten())
307 }
308}
309
310#[derive(Debug)]
311pub struct OpenRequest {
312 pub open_data: OpenData,
313 pub incoming_event: Option<Event>,
314 pub use_vtl2_connection_id: bool,
315}
316
317#[derive(Debug)]
318pub struct RestoreRequest {
319 pub incoming_event: Option<Event>,
320 pub redirected_event_flag: Option<u16>,
322 pub connection_id: u32,
324}
325
326pub enum ChannelRequest {
328 Open(FailableRpc<OpenRequest, OpenOutput>),
329 Restore(FailableRpc<RestoreRequest, OpenOutput>),
330 Close(Rpc<(), ()>),
331 Gpadl(FailableRpc<GpadlRequest, ()>),
332 TeardownGpadl(Rpc<GpadlId, ()>),
333 Modify(Rpc<ModifyRequest, i32>),
334}
335
336#[derive(Debug)]
337pub struct OpenOutput {
338 pub redirected_event_flag: Option<u16>,
340}
341
342impl std::fmt::Display for ChannelRequest {
343 fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
344 let s = match self {
345 ChannelRequest::Open(_) => "Open",
346 ChannelRequest::Close(_) => "Close",
347 ChannelRequest::Restore(_) => "Restore",
348 ChannelRequest::Gpadl(_) => "Gpadl",
349 ChannelRequest::TeardownGpadl(_) => "TeardownGpadl",
350 ChannelRequest::Modify(_) => "Modify",
351 };
352 fmt.pad(s)
353 }
354}
355
356#[derive(Debug, Error)]
357pub enum RestoreError {
358 #[error("unsupported protocol version {0:#x}")]
359 UnsupportedVersion(u32),
360
361 #[error("unsupported feature flags {0:#x}")]
362 UnsupportedFeatureFlags(u32),
363
364 #[error("duplicate channel id {0}")]
365 DuplicateChannelId(u32),
366
367 #[error("duplicate gpadl id {0}")]
368 DuplicateGpadlId(u32),
369
370 #[error("gpadl for unknown channel id {0}")]
371 GpadlForUnknownChannelId(u32),
372
373 #[error("invalid pending message")]
374 InvalidPendingMessage(#[source] vmbus_core::MessageTooLarge),
375
376 #[error("failed to offer restored channel")]
377 OfferFailed(#[source] anyhow::Error),
378}
379
380#[derive(Debug, Inspect)]
383pub struct OfferInfo {
384 pub offer: protocol::OfferChannel,
385 #[inspect(skip)]
386 pub guest_to_host_interrupt: Interrupt,
387 #[inspect(skip)]
388 pub request_send: mesh::Sender<ChannelRequest>,
389 #[inspect(skip)]
390 pub revoke_recv: mesh::OneshotReceiver<()>,
391}
392
393#[derive(Debug)]
394enum ClientRequest {
395 Connect(Rpc<ConnectRequest, Result<ConnectResult, ConnectError>>),
396 Unload(Rpc<(), ()>),
397 Modify(Rpc<ModifyConnectionRequest, ConnectionState>),
398 HvsockConnect(Rpc<HvsockConnectRequest, Option<OfferInfo>>),
399}
400
401impl std::fmt::Display for ClientRequest {
402 fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
403 let s = match self {
404 ClientRequest::Connect(..) => "Connect",
405 ClientRequest::Unload { .. } => "Unload",
406 ClientRequest::Modify(..) => "Modify",
407 ClientRequest::HvsockConnect(..) => "HvsockConnect",
408 };
409 fmt.pad(s)
410 }
411}
412
413enum TaskRequest {
414 Inspect(inspect::Deferred),
415 Save(Rpc<(), SavedState>),
416 Restore(Rpc<SavedState, Result<Option<ConnectResult>, RestoreError>>),
417 PostRestore(Rpc<(), ()>),
418 Start,
419 Stop(Rpc<(), ()>),
420}
421
422#[derive(Inspect)]
426#[inspect(external_tag)]
427enum ClientState {
428 Disconnected,
430 Connecting {
432 version: Version,
433 #[inspect(skip)]
434 rpc: Rpc<ConnectRequest, Result<ConnectResult, ConnectError>>,
435 },
436 Connected {
438 version: VersionInfo,
439 #[inspect(skip)]
440 offer_send: mesh::Sender<OfferInfo>,
441 },
442 RequestingOffers {
444 version: VersionInfo,
445 #[inspect(skip)]
446 rpc: Rpc<(), Result<ConnectResult, ConnectError>>,
447 #[inspect(skip)]
448 offers: Vec<OfferInfo>,
449 },
450 Disconnecting {
452 version: VersionInfo,
453 #[inspect(skip)]
454 rpc: Rpc<(), ()>,
455 },
456}
457
458impl ClientState {
459 fn get_version(&self) -> Option<VersionInfo> {
460 match self {
461 ClientState::Connected { version, .. } => Some(*version),
462 ClientState::RequestingOffers { version, .. } => Some(*version),
463 ClientState::Disconnecting { version, .. } => Some(*version),
464 ClientState::Disconnected | ClientState::Connecting { .. } => None,
465 }
466 }
467}
468
469impl std::fmt::Display for ClientState {
470 fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
471 let s = match self {
472 ClientState::Disconnected => "Disconnected",
473 ClientState::Connecting { .. } => "Connecting",
474 ClientState::Connected { .. } => "Connected",
475 ClientState::RequestingOffers { .. } => "RequestingOffers",
476 ClientState::Disconnecting { .. } => "Disconnecting",
477 };
478 fmt.pad(s)
479 }
480}
481
482#[derive(Copy, Clone, Debug, Default)]
483struct ConnectRequest {
484 target_message_vp: u32,
485 monitor_page: Option<MonitorPageGpas>,
486 client_id: Guid,
487}
488
489#[derive(Copy, Clone, Debug, Default)]
490pub struct ModifyConnectionRequest {
491 pub monitor_page: Option<MonitorPageGpas>,
492}
493
494impl From<ModifyConnectionRequest> for protocol::ModifyConnection {
495 fn from(value: ModifyConnectionRequest) -> Self {
496 let monitor_page = value.monitor_page.unwrap_or_default();
497
498 Self {
499 parent_to_child_monitor_page_gpa: monitor_page.parent_to_child,
500 child_to_parent_monitor_page_gpa: monitor_page.child_to_parent,
501 }
502 }
503}
504
505#[derive(Debug, Inspect)]
509#[inspect(external_tag)]
510enum ChannelState {
511 Offered,
513 Opening {
515 redirected_event_flag: Option<u16>,
516 #[inspect(skip)]
517 redirected_event: Option<Event>,
518 #[inspect(skip)]
519 rpc: FailableRpc<(), OpenOutput>,
520 },
521 Restored,
523 Opened {
525 redirected_event_flag: Option<u16>,
526 #[inspect(skip)]
527 redirected_event: Option<Event>,
528 },
529 Revoked,
531}
532
533impl std::fmt::Display for ChannelState {
534 fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
535 let s = match self {
536 ChannelState::Opening { .. } => "Opening",
537 ChannelState::Offered => "Offered",
538 ChannelState::Opened { .. } => "Opened",
539 ChannelState::Restored => "Restored",
540 ChannelState::Revoked => "Revoked",
541 };
542 fmt.pad(s)
543 }
544}
545
546#[derive(Debug, Inspect)]
547struct Channel {
548 offer: protocol::OfferChannel,
549 #[inspect(skip)]
551 revoke_send: Option<mesh::OneshotSender<()>>,
552 state: ChannelState,
553 #[inspect(with = "|x| x.is_some()")]
554 modify_response_send: Option<Rpc<(), i32>>,
555 #[inspect(with = "|x| inspect::iter_by_key(x).map_key(|x| x.0)")]
556 gpadls: HashMap<GpadlId, GpadlState>,
557 is_client_released: bool,
558 connection_id: Arc<AtomicU32>,
559}
560
561impl Channel {
562 fn pending_request(&self) -> Option<&'static str> {
563 if self.modify_response_send.is_some() {
564 return Some("modify");
565 }
566 self.gpadls.iter().find_map(|(_, gpadl)| match gpadl {
567 GpadlState::Offered(_) => Some("creating gpadl"),
568 GpadlState::Created => None,
569 GpadlState::TearingDown { .. } => Some("tearing down gpadl"),
570 })
571 }
572}
573
574#[derive(Inspect)]
575struct ClientTask {
576 #[inspect(flatten)]
577 inner: ClientTaskInner,
578 channels: ChannelList,
579 state: ClientState,
580 hvsock_tracker: hvsock::HvsockRequestTracker,
581 running: bool,
582 #[inspect(with = "|x| x.is_some()")]
583 modify_request: Option<Rpc<ModifyConnectionRequest, ConnectionState>>,
584 #[inspect(skip)]
585 msg_source: Box<dyn VmbusMessageSource>,
586 #[inspect(skip)]
587 task_recv: mesh::Receiver<TaskRequest>,
588 #[inspect(skip)]
589 client_request_recv: mesh::Receiver<ClientRequest>,
590}
591
592impl ClientTask {
593 fn handle_initiate_contact(
594 &mut self,
595 rpc: Rpc<ConnectRequest, Result<ConnectResult, ConnectError>>,
596 version: Version,
597 ) {
598 let ClientState::Disconnected = self.state else {
599 tracing::warn!(client_state = %self.state, "invalid client state for InitiateContact");
600 rpc.complete(Err(ConnectError::InvalidState));
601 return;
602 };
603 let feature_flags = if version >= Version::Copper {
604 SUPPORTED_FEATURE_FLAGS
605 } else {
606 FeatureFlags::new()
607 };
608
609 let request = rpc.input();
610
611 tracing::debug!(version = ?version, ?feature_flags, "VmBus client connecting");
612 let target_info = protocol::TargetInfo::new()
613 .with_sint(SINT)
614 .with_vtl(VTL)
615 .with_feature_flags(feature_flags.into());
616 let monitor_page = request.monitor_page.unwrap_or_default();
617 let msg = protocol::InitiateContact2 {
618 initiate_contact: protocol::InitiateContact {
619 version_requested: version as u32,
620 target_message_vp: request.target_message_vp,
621 interrupt_page_or_target_info: target_info.into(),
622 parent_to_child_monitor_page_gpa: monitor_page.parent_to_child,
623 child_to_parent_monitor_page_gpa: monitor_page.child_to_parent,
624 },
625 client_id: request.client_id,
626 };
627
628 self.state = ClientState::Connecting { version, rpc };
629 if version < Version::Copper {
630 self.inner.messages.send(&msg.initiate_contact)
631 } else {
632 self.inner.messages.send(&msg);
633 }
634 }
635
636 fn handle_unload(&mut self, rpc: Rpc<(), ()>) {
637 tracing::debug!(%self.state, "VmBus client disconnecting");
638 self.state = ClientState::Disconnecting {
639 version: self.state.get_version().expect("invalid state for unload"),
640 rpc,
641 };
642
643 self.inner.messages.send(&protocol::Unload {});
644 }
645
646 fn handle_modify(&mut self, request: Rpc<ModifyConnectionRequest, ConnectionState>) {
647 if !matches!(self.state, ClientState::Connected { version, .. }
648 if version.feature_flags.modify_connection())
649 {
650 tracing::warn!("ModifyConnection not supported");
651 request.complete(ConnectionState::FAILED_UNKNOWN_FAILURE);
652 return;
653 }
654
655 if self.modify_request.is_some() {
656 tracing::warn!("Duplicate ModifyConnection request");
657 request.complete(ConnectionState::FAILED_UNKNOWN_FAILURE);
658 return;
659 }
660
661 let message = protocol::ModifyConnection::from(*request.input());
662 self.modify_request = Some(request);
663 self.inner.messages.send(&message);
664 }
665
666 fn handle_tl_connect(&mut self, rpc: Rpc<HvsockConnectRequest, Option<OfferInfo>>) {
667 let message = protocol::TlConnectRequest2::from(*rpc.input());
671 self.hvsock_tracker.add_request(rpc);
672 self.inner.messages.send(&message);
673 }
674
675 fn handle_client_request(&mut self, request: ClientRequest) {
676 match request {
677 ClientRequest::Connect(rpc) => {
678 self.handle_initiate_contact(rpc, *SUPPORTED_VERSIONS.last().unwrap());
679 }
680 ClientRequest::Unload(rpc) => {
681 self.handle_unload(rpc);
682 }
683 ClientRequest::Modify(request) => self.handle_modify(request),
684 ClientRequest::HvsockConnect(request) => self.handle_tl_connect(request),
685 }
686 }
687
688 fn handle_version_response(&mut self, msg: protocol::VersionResponse2) {
689 let old_state = std::mem::replace(&mut self.state, ClientState::Disconnected);
690 let ClientState::Connecting { version, rpc } = old_state else {
691 self.state = old_state;
692 tracing::warn!(
693 client_state = %self.state,
694 "invalid client state to handle VersionResponse"
695 );
696 return;
697 };
698 if msg.version_response.version_supported > 0 {
699 if msg.version_response.connection_state != ConnectionState::SUCCESSFUL {
700 rpc.complete(Err(ConnectError::FailedToConnect(
701 msg.version_response.connection_state,
702 )));
703 return;
704 }
705
706 let feature_flags = if version >= Version::Copper {
707 FeatureFlags::from(msg.supported_features)
708 } else {
709 FeatureFlags::new()
710 };
711
712 let version = VersionInfo {
713 version,
714 feature_flags,
715 };
716
717 self.inner.messages.send(&protocol::RequestOffers {});
718 self.state = ClientState::RequestingOffers {
719 version,
720 rpc: rpc.split().1,
721 offers: Vec::new(),
722 };
723 tracing::info!(?version, "VmBus client connected, requesting offers");
724 } else {
725 let index = SUPPORTED_VERSIONS
726 .iter()
727 .position(|v| *v == version)
728 .unwrap();
729
730 if index == 0 {
731 rpc.complete(Err(ConnectError::NoSupportedVersions));
732 return;
733 }
734 let next_version = SUPPORTED_VERSIONS[index - 1];
735 tracing::debug!(
736 version = version as u32,
737 next_version = next_version as u32,
738 "Unsupported version, retrying"
739 );
740 self.handle_initiate_contact(rpc, next_version);
741 }
742 }
743
744 fn create_channel(&mut self, offer: protocol::OfferChannel) -> Result<OfferInfo> {
745 self.create_channel_core(offer, ChannelState::Offered)
746 }
747
748 fn create_channel_core(
749 &mut self,
750 offer: protocol::OfferChannel,
751 state: ChannelState,
752 ) -> Result<OfferInfo> {
753 if self.channels.0.contains_key(&offer.channel_id) {
754 anyhow::bail!("channel {:?} exists", offer.channel_id);
755 }
756 let (request_send, request_recv) = mesh::channel();
757 let (revoke_send, revoke_recv) = mesh::oneshot();
758
759 let connection_id = Arc::new(AtomicU32::new(0));
760 self.channels.0.insert(
761 offer.channel_id,
762 Channel {
763 revoke_send: Some(revoke_send),
764 offer,
765 state,
766 modify_response_send: None,
767 gpadls: HashMap::new(),
768 is_client_released: false,
769 connection_id: connection_id.clone(),
770 },
771 );
772
773 self.inner
774 .channel_requests
775 .push(TaggedStream::new(offer.channel_id, request_recv));
776
777 Ok(OfferInfo {
778 offer,
779 guest_to_host_interrupt: self.inner.synic.guest_to_host_interrupt(connection_id),
780 revoke_recv,
781 request_send,
782 })
783 }
784
785 fn handle_offer(&mut self, offer: protocol::OfferChannel) {
786 let offer_info = self
787 .create_channel(offer)
788 .expect("channel should not exist");
789
790 tracing::info!(
791 state = %self.state,
792 channel_id = offer.channel_id.0,
793 interface_id = %offer.interface_id,
794 instance_id = %offer.instance_id,
795 subchannel_index = offer.subchannel_index,
796 "received offer");
797
798 if let Some(offer) = self.hvsock_tracker.check_offer(&offer_info.offer) {
799 offer.complete(Some(offer_info));
800 } else {
801 match &mut self.state {
802 ClientState::Connected { offer_send, .. } => {
803 offer_send.send(offer_info);
804 }
805 ClientState::RequestingOffers { offers, .. } => {
806 offers.push(offer_info);
807 }
808 state => unreachable!("invalid client state for OfferChannel: {state}"),
809 }
810 }
811 }
812
813 fn handle_rescind(&mut self, rescind: protocol::RescindChannelOffer) -> TriedRelease {
814 let mut channel = self.channels.get_mut(rescind.channel_id);
815 tracing::info!(
816 state = %self.state,
817 channel_id = rescind.channel_id.0,
818 key = %OfferKey::from(&channel.offer),
819 "received rescind"
820 );
821 let event_flag = match std::mem::replace(&mut channel.state, ChannelState::Revoked) {
822 ChannelState::Offered => None,
823 ChannelState::Opening {
824 redirected_event_flag,
825 redirected_event: _,
826 rpc,
827 } => {
828 rpc.fail(anyhow::anyhow!("channel revoked"));
829 redirected_event_flag
830 }
831 ChannelState::Restored => None,
832 ChannelState::Opened {
833 redirected_event_flag,
834 redirected_event: _,
835 } => redirected_event_flag,
836 ChannelState::Revoked => {
837 panic!("channel id {:?} already revoked", rescind.channel_id);
838 }
839 };
840 if let Some(event_flag) = event_flag {
841 self.inner.synic.free_event_flag(event_flag);
842 }
843
844 channel.revoke_send.take().unwrap().send(());
846
847 channel.try_release(&mut self.inner.messages)
848 }
849
850 fn handle_offers_delivered(&mut self) {
851 match std::mem::replace(&mut self.state, ClientState::Disconnected) {
852 ClientState::RequestingOffers {
853 version,
854 rpc,
855 offers,
856 } => {
857 tracing::info!(version = ?version, "VmBus client connected, offers delivered");
858 let (offer_send, offer_recv) = mesh::channel();
859 self.state = ClientState::Connected {
860 version,
861 offer_send,
862 };
863 rpc.complete(Ok(ConnectResult {
864 version,
865 offers,
866 offer_recv,
867 }));
868 }
869 state => {
870 tracing::warn!(client_state = %state, "invalid client state for OffersDelivered");
871 self.state = state;
872 }
873 }
874 }
875
876 fn handle_gpadl_created(&mut self, request: protocol::GpadlCreated) -> TriedRelease {
877 let mut channel = self.channels.get_mut(request.channel_id);
878 let Some(gpadl_state) = channel.gpadls.get_mut(&request.gpadl_id) else {
879 panic!("GpadlCreated for unknown gpadl {:#x}", request.gpadl_id.0);
880 };
881
882 let rpc = match std::mem::replace(gpadl_state, GpadlState::Created) {
883 GpadlState::Offered(rpc) => rpc,
884 old_state => {
885 panic!(
886 "invalid state {old_state:?} for gpadl {:#x}:{:#x}",
887 request.channel_id.0, request.gpadl_id.0
888 );
889 }
890 };
891
892 let gpadl_created = request.status == protocol::STATUS_SUCCESS;
893 if gpadl_created {
894 rpc.complete(Ok(()));
895 } else {
896 channel.gpadls.remove(&request.gpadl_id).unwrap();
897 rpc.fail(anyhow::anyhow!(
898 "gpadl creation failed: {:#x}",
899 request.status
900 ));
901 };
902 channel.try_release(&mut self.inner.messages)
903 }
904
905 fn handle_open_result(&mut self, result: protocol::OpenResult) {
906 let mut channel = self.channels.get_mut(result.channel_id);
907 tracing::debug!(
908 channel_id = result.channel_id.0,
909 key = %OfferKey::from(&channel.offer),
910 result = result.status,
911 "received open result"
912 );
913
914 let channel_opened = result.status == protocol::STATUS_SUCCESS as u32;
915 let old_state = std::mem::replace(&mut channel.state, ChannelState::Offered);
916 let ChannelState::Opening {
917 redirected_event_flag,
918 redirected_event,
919 rpc,
920 } = old_state
921 else {
922 tracing::warn!(
923 key = %OfferKey::from(&channel.offer),
924 old_state = ?channel.state,
925 channel_opened,
926 "invalid state for open result"
927 );
928 channel.state = old_state;
929 return;
930 };
931
932 if !channel_opened {
933 if let Some(event_flag) = redirected_event_flag {
934 self.inner.synic.free_event_flag(event_flag);
935 }
936 rpc.fail(anyhow::anyhow!("open failed: {:#x}", result.status));
937 return;
938 }
939
940 channel.state = ChannelState::Opened {
941 redirected_event_flag,
942 redirected_event,
943 };
944
945 rpc.complete(Ok(OpenOutput {
946 redirected_event_flag,
947 }));
948 }
949
950 fn handle_gpadl_torndown(&mut self, request: protocol::GpadlTorndown) -> TriedRelease {
951 let Some(channel_id) = self.inner.teardown_gpadls.remove(&request.gpadl_id) else {
952 panic!("gpadl {:#x} not in teardown list", request.gpadl_id.0);
953 };
954
955 let mut channel = self.channels.get_mut(channel_id);
956 tracing::debug!(
957 gpadl_id = request.gpadl_id.0,
958 channel_id = channel_id.0,
959 key = %OfferKey::from(&channel.offer),
960 "Received GpadlTorndown"
961 );
962
963 let gpadl_state = channel
964 .gpadls
965 .remove(&request.gpadl_id)
966 .expect("gpadl validated above");
967
968 let GpadlState::TearingDown { rpcs } = gpadl_state else {
969 panic!("gpadl should be tearing down if in teardown list, state = {gpadl_state:?}");
970 };
971
972 for rpc in rpcs {
973 rpc.complete(());
974 }
975 channel.try_release(&mut self.inner.messages)
976 }
977
978 fn handle_unload_complete(&mut self) {
979 match std::mem::replace(&mut self.state, ClientState::Disconnected) {
980 ClientState::Disconnecting { version: _, rpc } => {
981 tracing::info!("VmBus client disconnected");
982 rpc.complete(());
983 }
984 state => {
985 tracing::warn!(client_state = %state, "invalid client state for UnloadComplete");
986 }
987 }
988 }
989
990 fn handle_modify_complete(&mut self, response: protocol::ModifyConnectionResponse) {
991 if let Some(request) = self.modify_request.take() {
992 request.complete(response.connection_state)
993 } else {
994 tracing::warn!("Unexpected modify complete request");
995 }
996 }
997
998 fn handle_modify_channel_response(
999 &mut self,
1000 response: protocol::ModifyChannelResponse,
1001 ) -> TriedRelease {
1002 let mut channel = self.channels.get_mut(response.channel_id);
1003 let Some(sender) = channel.modify_response_send.take() else {
1004 panic!(
1005 "unexpected modify channel response for channel {:#x}",
1006 response.channel_id.0
1007 );
1008 };
1009
1010 sender.complete(response.status);
1011 channel.try_release(&mut self.inner.messages)
1012 }
1013
1014 fn handle_tl_connect_result(&mut self, response: protocol::TlConnectResult) {
1015 if let Some(rpc) = self.hvsock_tracker.check_result(&response) {
1016 rpc.complete(None);
1017 }
1018 }
1019
1020 fn handle_synic_message(&mut self, data: &[u8]) -> bool {
1022 let msg = Message::parse(data, self.state.get_version()).unwrap();
1023 tracing::trace!(?msg, "received client message from synic");
1024
1025 match msg {
1026 Message::VersionResponse3(version_response, ..) => {
1027 self.handle_version_response(version_response.version_response2);
1032 }
1033 Message::VersionResponse2(version_response, ..) => {
1034 self.handle_version_response(version_response);
1035 }
1036 Message::VersionResponse(version_response, ..) => {
1037 self.handle_version_response(version_response.into());
1038 }
1039 Message::OfferChannel(offer, ..) => {
1040 self.handle_offer(offer);
1041 }
1042 Message::AllOffersDelivered(..) => {
1043 self.handle_offers_delivered();
1044 }
1045 Message::UnloadComplete(..) => {
1046 self.handle_unload_complete();
1047 }
1048 Message::ModifyConnectionResponse(response, ..) => {
1049 self.handle_modify_complete(response);
1050 }
1051 Message::GpadlCreated(gpadl, ..) => {
1052 self.handle_gpadl_created(gpadl);
1053 }
1054 Message::OpenResult(result, ..) => {
1055 self.handle_open_result(result);
1056 }
1057 Message::GpadlTorndown(gpadl, ..) => {
1058 self.handle_gpadl_torndown(gpadl);
1059 }
1060 Message::RescindChannelOffer(rescind, ..) => {
1061 self.handle_rescind(rescind);
1062 }
1063 Message::ModifyChannelResponse(response, ..) => {
1064 self.handle_modify_channel_response(response);
1065 }
1066 Message::TlConnectResult(response, ..) => self.handle_tl_connect_result(response),
1067 Message::CloseReservedChannelResponse(..) => {
1069 todo!("Unsupported message {msg:?}")
1070 }
1071 Message::PauseResponse(..) => {
1072 return false;
1073 }
1074 Message::RequestOffers(..)
1076 | Message::OpenChannel2(..)
1077 | Message::OpenChannel(..)
1078 | Message::CloseChannel(..)
1079 | Message::GpadlHeader(..)
1080 | Message::GpadlBody(..)
1081 | Message::GpadlTeardown(..)
1082 | Message::RelIdReleased(..)
1083 | Message::InitiateContact(..)
1084 | Message::InitiateContact2(..)
1085 | Message::Unload(..)
1086 | Message::OpenReservedChannel(..)
1087 | Message::CloseReservedChannel(..)
1088 | Message::TlConnectRequest2(..)
1089 | Message::TlConnectRequest(..)
1090 | Message::ModifyChannel(..)
1091 | Message::ModifyConnection(..)
1092 | Message::Pause(..)
1093 | Message::Resume(..) => {
1094 unreachable!("Client received server message {msg:?}");
1095 }
1096 }
1097 true
1098 }
1099
1100 fn handle_open_channel(
1101 &mut self,
1102 channel_id: ChannelId,
1103 rpc: FailableRpc<OpenRequest, OpenOutput>,
1104 ) {
1105 let mut channel = self.channels.get_mut(channel_id);
1106 match &channel.state {
1107 ChannelState::Offered => {}
1108 ChannelState::Revoked => {
1109 rpc.fail(anyhow::anyhow!("channel revoked"));
1110 return;
1111 }
1112 state => {
1113 rpc.fail(anyhow::anyhow!("invalid channel state: {}", state));
1114 return;
1115 }
1116 }
1117
1118 tracing::info!(
1119 channel_id = channel_id.0,
1120 key = %OfferKey::from(&channel.offer),
1121 "opening channel on host"
1122 );
1123
1124 let (request, rpc) = rpc.split();
1125 let open_data = &request.open_data;
1126
1127 let supports_interrupt_redirection =
1128 if let ClientState::Connected { version, .. } = self.state {
1129 version.feature_flags.guest_specified_signal_parameters()
1130 || version.feature_flags.channel_interrupt_redirection()
1131 } else {
1132 false
1133 };
1134
1135 if !supports_interrupt_redirection && open_data.event_flag != channel_id.0 as u16 {
1136 rpc.fail(anyhow::anyhow!(
1137 "host does not support specifying the event flag"
1138 ));
1139 return;
1140 }
1141
1142 let open_channel = protocol::OpenChannel {
1143 channel_id,
1144 open_id: 0,
1145 ring_buffer_gpadl_id: open_data.ring_gpadl_id,
1146 target_vp: open_data
1147 .target_vp
1148 .unwrap_or(protocol::VP_INDEX_DISABLE_INTERRUPT),
1149 downstream_ring_buffer_page_offset: open_data.ring_offset,
1150 user_data: open_data.user_data,
1151 };
1152
1153 let connection_id = if request.use_vtl2_connection_id {
1154 if !supports_interrupt_redirection {
1155 rpc.fail(anyhow::anyhow!(
1156 "host does not support specfiying the connection ID"
1157 ));
1158 return;
1159 }
1160 protocol::ConnectionId::new(channel_id.0, 2.try_into().unwrap(), 7).0
1161 } else {
1162 open_data.connection_id
1163 };
1164
1165 let mut flags = OpenChannelFlags::new();
1168 let event_flag = if let Some(event) = &request.incoming_event {
1169 if !supports_interrupt_redirection {
1170 rpc.fail(anyhow::anyhow!(
1171 "host does not support redirecting interrupts"
1172 ));
1173 return;
1174 }
1175
1176 flags.set_redirect_interrupt(true);
1177 match self.inner.synic.allocate_event_flag(event) {
1178 Ok(flag) => flag,
1179 Err(err) => {
1180 rpc.fail(err.context("failed to allocate event flag"));
1181 return;
1182 }
1183 }
1184 } else {
1185 open_data.event_flag
1186 };
1187
1188 if supports_interrupt_redirection {
1189 self.inner.messages.send(&protocol::OpenChannel2 {
1190 open_channel,
1191 connection_id,
1192 event_flag,
1193 flags,
1194 });
1195 } else {
1196 self.inner.messages.send(&open_channel);
1197 }
1198
1199 channel
1200 .connection_id
1201 .store(connection_id, Ordering::Release);
1202 channel.state = ChannelState::Opening {
1203 redirected_event_flag: (request.incoming_event.is_some()).then_some(event_flag),
1204 redirected_event: request.incoming_event,
1205 rpc,
1206 }
1207 }
1208
1209 fn handle_restore_channel(
1210 &mut self,
1211 channel_id: ChannelId,
1212 request: RestoreRequest,
1213 ) -> Result<OpenOutput> {
1214 let mut channel = self.channels.get_mut(channel_id);
1215 if !matches!(channel.state, ChannelState::Restored) {
1216 anyhow::bail!("invalid channel state: {}", channel.state);
1217 }
1218
1219 if request.incoming_event.is_some() != request.redirected_event_flag.is_some() {
1220 anyhow::bail!("incoming event and redirected event flag must both be set or unset");
1221 }
1222
1223 if let Some((flag, event)) = request
1224 .redirected_event_flag
1225 .zip(request.incoming_event.as_ref())
1226 {
1227 self.inner.synic.restore_event_flag(flag, event)?;
1228 }
1229
1230 channel
1231 .connection_id
1232 .store(request.connection_id, Ordering::Release);
1233 channel.state = ChannelState::Opened {
1234 redirected_event_flag: request.redirected_event_flag,
1235 redirected_event: request.incoming_event,
1236 };
1237 Ok(OpenOutput {
1238 redirected_event_flag: request.redirected_event_flag,
1239 })
1240 }
1241
1242 fn handle_gpadl(&mut self, channel_id: ChannelId, rpc: FailableRpc<GpadlRequest, ()>) {
1243 let (request, rpc) = rpc.split();
1244 let mut channel = self.channels.get_mut(channel_id);
1245 if channel
1246 .gpadls
1247 .insert(request.id, GpadlState::Offered(rpc))
1248 .is_some()
1249 {
1250 panic!(
1251 "duplicate gpadl ID {:?} for channel {:?}.",
1252 request.id, channel_id
1253 );
1254 }
1255
1256 tracing::trace!(
1257 channel_id = channel_id.0,
1258 key = %OfferKey::from(&channel.offer),
1259 gpadl_id = request.id.0,
1260 count = request.count,
1261 len = request.buf.len(),
1262 "received gpadl request"
1263 );
1264
1265 let (first, remaining) = if request.buf.len() > protocol::GpadlHeader::MAX_DATA_VALUES {
1267 request.buf.split_at(protocol::GpadlHeader::MAX_DATA_VALUES)
1268 } else {
1269 (request.buf.as_slice(), [].as_slice())
1270 };
1271
1272 let message = protocol::GpadlHeader {
1273 channel_id,
1274 gpadl_id: request.id,
1275 len: (request.buf.len() * size_of::<u64>())
1276 .try_into()
1277 .expect("Too many GPA values"),
1278 count: request.count,
1279 };
1280
1281 self.inner
1282 .messages
1283 .send_with_data(&message, first.as_bytes());
1284
1285 let message = protocol::GpadlBody {
1287 rsvd: 0,
1288 gpadl_id: request.id,
1289 };
1290 for chunk in remaining.chunks(protocol::GpadlBody::MAX_DATA_VALUES) {
1291 self.inner
1292 .messages
1293 .send_with_data(&message, chunk.as_bytes());
1294 }
1295 }
1296
1297 fn handle_gpadl_teardown(&mut self, channel_id: ChannelId, rpc: Rpc<GpadlId, ()>) {
1298 let (gpadl_id, rpc) = rpc.split();
1299 let mut channel = self.channels.get_mut(channel_id);
1300 let Some(gpadl_state) = channel.gpadls.get_mut(&gpadl_id) else {
1301 tracing::warn!(
1302 gpadl_id = gpadl_id.0,
1303 channel_id = channel_id.0,
1304 key = %OfferKey::from(&channel.offer),
1305 "Gpadl teardown for unknown gpadl or revoked channel"
1306 );
1307 return;
1308 };
1309
1310 match gpadl_state {
1311 GpadlState::Offered(_) => {
1312 tracing::warn!(
1313 gpadl_id = gpadl_id.0,
1314 channel_id = channel_id.0,
1315 key = %OfferKey::from(&channel.offer),
1316 "gpadl teardown for offered gpadl"
1317 );
1318 }
1319 GpadlState::Created => {
1320 *gpadl_state = GpadlState::TearingDown { rpcs: vec![rpc] };
1321 assert!(
1325 self.inner
1326 .teardown_gpadls
1327 .insert(gpadl_id, channel_id)
1328 .is_none(),
1329 "Gpadl state validated above"
1330 );
1331
1332 self.inner.messages.send(&protocol::GpadlTeardown {
1333 channel_id,
1334 gpadl_id,
1335 });
1336 }
1337 GpadlState::TearingDown { rpcs } => {
1338 rpcs.push(rpc);
1339 }
1340 }
1341 }
1342
1343 fn handle_close_channel(&mut self, channel_id: ChannelId) {
1344 let mut channel = self.channels.get_mut(channel_id);
1345 self.inner.close_channel(channel_id, &mut channel);
1346 }
1347
1348 fn handle_modify_channel(&mut self, channel_id: ChannelId, rpc: Rpc<ModifyRequest, i32>) {
1349 assert!(self.check_version(Version::Iron));
1353 let mut channel = self.channels.get_mut(channel_id);
1354 if channel.modify_response_send.is_some() {
1355 panic!("duplicate channel modify request {channel_id:?}");
1356 }
1357
1358 let (request, response) = rpc.split();
1359 channel.modify_response_send = Some(response);
1360 let payload = match request {
1361 ModifyRequest::TargetVp { target_vp } => protocol::ModifyChannel {
1362 channel_id,
1363 target_vp,
1364 },
1365 };
1366
1367 self.inner.messages.send(&payload);
1368 }
1369
1370 fn handle_channel_request(&mut self, channel_id: ChannelId, request: ChannelRequest) {
1371 match request {
1372 ChannelRequest::Open(rpc) => self.handle_open_channel(channel_id, rpc),
1373 ChannelRequest::Restore(rpc) => {
1374 rpc.handle_failable_sync(|request| self.handle_restore_channel(channel_id, request))
1375 }
1376 ChannelRequest::Gpadl(req) => self.handle_gpadl(channel_id, req),
1377 ChannelRequest::TeardownGpadl(req) => self.handle_gpadl_teardown(channel_id, req),
1378 ChannelRequest::Close(req) => {
1379 req.handle_sync(|()| self.handle_close_channel(channel_id))
1380 }
1381 ChannelRequest::Modify(req) => self.handle_modify_channel(channel_id, req),
1382 }
1383 }
1384
1385 async fn handle_task(&mut self, task: TaskRequest) {
1386 match task {
1387 TaskRequest::Inspect(deferred) => {
1388 deferred.inspect(&*self);
1389 }
1390 TaskRequest::Save(rpc) => rpc.handle_sync(|()| self.handle_save()),
1391 TaskRequest::Restore(rpc) => {
1392 rpc.handle_sync(|saved_state| self.handle_restore(saved_state))
1393 }
1394 TaskRequest::PostRestore(rpc) => rpc.handle_sync(|()| self.handle_post_restore()),
1395 TaskRequest::Start => self.handle_start(),
1396 TaskRequest::Stop(rpc) => rpc.handle(async |()| self.handle_stop().await).await,
1397 }
1398 }
1399
1400 fn handle_device_removal(&mut self, channel_id: ChannelId) -> TriedRelease {
1402 let mut channel = self.channels.get_mut(channel_id);
1403 channel.is_client_released = true;
1404 if let ChannelState::Opened { .. } = channel.state {
1406 tracing::warn!(
1407 channel_id = channel_id.0,
1408 key = %OfferKey::from(&channel.offer),
1409 "Channel dropped without closing first"
1410 );
1411 self.inner.close_channel(channel_id, &mut channel);
1412 }
1413 channel.try_release(&mut self.inner.messages)
1414 }
1415
1416 fn check_version(&self, version: Version) -> bool {
1418 matches!(self.state, ClientState::Connected { version: v, .. } if v.version >= version)
1419 }
1420
1421 fn handle_start(&mut self) {
1422 assert!(!self.running);
1423 self.msg_source.resume_message_stream();
1424 self.inner.messages.resume();
1425 self.running = true;
1426 }
1427
1428 async fn handle_stop(&mut self) {
1429 assert!(self.running);
1430
1431 loop {
1432 while let Some((id, request)) = self.channels.revoked_channel_with_pending_request() {
1437 tracelimit::info_ratelimited!(
1438 channel_id = id.0,
1439 request,
1440 "waiting for responses for channel"
1441 );
1442 assert!(self.process_next_message().await);
1443 }
1444
1445 if self.can_pause_resume() {
1446 self.inner.messages.pause();
1447 } else {
1448 self.msg_source.pause_message_stream();
1451 self.inner.messages.force_pause();
1452 }
1453
1454 while self.process_next_message().await {}
1457
1458 if self
1461 .channels
1462 .revoked_channel_with_pending_request()
1463 .is_none()
1464 {
1465 break;
1466 }
1467 if !self.can_pause_resume() {
1468 self.msg_source.resume_message_stream();
1469 }
1470 self.inner.messages.resume();
1471 }
1472
1473 tracing::debug!("messages drained");
1474 self.running = false;
1476 }
1477
1478 async fn process_next_message(&mut self) -> bool {
1479 let mut buf = [0; protocol::MAX_MESSAGE_SIZE];
1480 let recv = self.msg_source.recv(&mut buf);
1481 let flush = async {
1484 self.inner.messages.flush_messages().await;
1485 std::future::pending().await
1486 };
1487 let size = (recv, flush)
1488 .race()
1489 .await
1490 .expect("Fatal error reading messages from synic");
1491 if size == 0 {
1492 return false;
1493 }
1494 self.handle_synic_message(&buf[..size])
1495 }
1496
1497 fn can_pause_resume(&self) -> bool {
1506 if let ClientState::Connected { version, .. } = self.state {
1507 version.feature_flags.pause_resume()
1508 } else {
1509 false
1510 }
1511 }
1512
1513 async fn run(&mut self) {
1514 let mut buf = [0; protocol::MAX_MESSAGE_SIZE];
1515 loop {
1516 let mut message_recv =
1517 OptionFuture::from(self.running.then(|| self.msg_source.recv(&mut buf).fuse()));
1518
1519 let host_backed_up = !self.inner.messages.is_empty();
1529 let flush_messages = OptionFuture::from(
1530 (self.running && host_backed_up)
1531 .then(|| self.inner.messages.flush_messages().fuse()),
1532 );
1533
1534 let mut client_request_recv = OptionFuture::from(
1535 (self.running && !host_backed_up).then(|| self.client_request_recv.next()),
1536 );
1537
1538 let mut channel_requests = OptionFuture::from(
1539 (self.running && !host_backed_up)
1540 .then(|| self.inner.channel_requests.select_next_some()),
1541 );
1542
1543 futures::select! { _r = pin!(flush_messages) => {}
1545 r = self.task_recv.next() => {
1546 if let Some(task) = r {
1547 self.handle_task(task).await;
1548 } else {
1549 break;
1550 }
1551 }
1552 r = client_request_recv => {
1553 if let Some(Some(request)) = r {
1554 self.handle_client_request(request);
1555 } else {
1556 break;
1557 }
1558 }
1559 r = channel_requests => {
1560 match r.unwrap() {
1561 (id, Some(request)) => self.handle_channel_request(id, request),
1562 (id, _) => {
1563 self.handle_device_removal(id);
1564 }
1565 }
1566 }
1567 r = message_recv => {
1568 match r.unwrap() {
1569 Ok(size) => {
1570 if size == 0 {
1571 panic!("Unexpected end of file reading messages from synic.");
1572 }
1573
1574 self.handle_synic_message(&buf[..size]);
1575 }
1576 Err(err) => {
1577 panic!("Error reading messages from synic: {err:?}");
1578 }
1579 }
1580 }
1581 complete => break,
1582 }
1583 }
1584 }
1585}
1586
1587impl ClientTaskInner {
1588 fn close_channel(&mut self, channel_id: ChannelId, channel: &mut Channel) {
1589 if let ChannelState::Opened {
1590 redirected_event_flag,
1591 ..
1592 } = channel.state
1593 {
1594 if let Some(flag) = redirected_event_flag {
1595 self.synic.free_event_flag(flag);
1596 }
1597 tracing::info!(
1598 channel_id = channel_id.0,
1599 key = %OfferKey::from(&channel.offer),
1600 "closing channel on host"
1601 );
1602
1603 self.messages.send(&protocol::CloseChannel { channel_id });
1604 channel.state = ChannelState::Offered;
1605 channel.connection_id.store(0, Ordering::Release);
1606 } else if matches!(channel.state, ChannelState::Revoked) {
1607 tracing::debug!(
1608 channel_id = channel_id.0,
1609 key = %OfferKey::from(&channel.offer),
1610 "close for channel already revoked by the server"
1611 );
1612 } else {
1613 tracing::warn!(
1614 channel_id = channel_id.0,
1615 key = %OfferKey::from(&channel.offer),
1616 channel_state = %channel.state,
1617 "invalid channel state for close channel"
1618 );
1619 }
1620 }
1621}
1622
1623#[derive(Debug, Inspect)]
1624#[inspect(external_tag)]
1625enum GpadlState {
1626 Offered(#[inspect(skip)] FailableRpc<(), ()>),
1628 Created,
1630 TearingDown {
1632 #[inspect(skip)]
1633 rpcs: Vec<Rpc<(), ()>>,
1634 },
1635}
1636
1637#[derive(Inspect)]
1638struct OutgoingMessages {
1639 #[inspect(skip)]
1640 poster: Box<dyn PollPostMessage>,
1641 #[inspect(with = "|x| x.len()")]
1642 queued: VecDeque<OutgoingMessage>,
1643 state: OutgoingMessageState,
1644}
1645
1646#[derive(Inspect, PartialEq, Eq, Debug)]
1647enum OutgoingMessageState {
1648 Running,
1649 SendingPauseMessage,
1650 Paused,
1651}
1652
1653impl OutgoingMessages {
1654 fn send<T: IntoBytes + protocol::VmbusMessage + std::fmt::Debug + Immutable + KnownLayout>(
1655 &mut self,
1656 msg: &T,
1657 ) {
1658 self.send_with_data(msg, &[])
1659 }
1660
1661 fn send_with_data<
1662 T: IntoBytes + protocol::VmbusMessage + std::fmt::Debug + Immutable + KnownLayout,
1663 >(
1664 &mut self,
1665 msg: &T,
1666 data: &[u8],
1667 ) {
1668 tracing::trace!(typ = ?T::MESSAGE_TYPE, "Sending message to host");
1669 let msg = OutgoingMessage::with_data(msg, data);
1670 if self.queued.is_empty() && self.state == OutgoingMessageState::Running {
1671 let r = self.poster.poll_post_message(
1672 &mut Context::from_waker(std::task::Waker::noop()),
1673 protocol::VMBUS_MESSAGE_REDIRECT_CONNECTION_ID,
1674 1,
1675 msg.data(),
1676 );
1677 if let Poll::Ready(()) = r {
1678 return;
1679 }
1680 }
1681 tracing::trace!("queueing message");
1682 self.queued.push_back(msg);
1683 }
1684
1685 async fn flush_messages(&mut self) {
1686 let mut send = async |msg: &OutgoingMessage| {
1687 poll_fn(|cx| {
1688 self.poster.poll_post_message(
1689 cx,
1690 protocol::VMBUS_MESSAGE_REDIRECT_CONNECTION_ID,
1691 1,
1692 msg.data(),
1693 )
1694 })
1695 .await
1696 };
1697 match self.state {
1698 OutgoingMessageState::Running => {
1699 while let Some(msg) = self.queued.front() {
1700 send(msg).await;
1701 tracing::trace!("sent queued message");
1702 self.queued.pop_front();
1703 }
1704 }
1705 OutgoingMessageState::SendingPauseMessage => {
1706 send(&OutgoingMessage::new(&protocol::Pause)).await;
1707 tracing::trace!("sent pause message");
1708 self.state = OutgoingMessageState::Paused;
1709 }
1710 OutgoingMessageState::Paused => {}
1711 }
1712 }
1713
1714 fn pause(&mut self) {
1717 assert_eq!(self.state, OutgoingMessageState::Running);
1718 self.state = OutgoingMessageState::SendingPauseMessage;
1719 self.queued
1721 .push_front(OutgoingMessage::new(&protocol::Resume));
1722 }
1723
1724 fn force_pause(&mut self) {
1728 assert_eq!(self.state, OutgoingMessageState::Running);
1729 self.state = OutgoingMessageState::Paused;
1730 }
1731
1732 fn resume(&mut self) {
1733 assert_eq!(self.state, OutgoingMessageState::Paused);
1734 self.state = OutgoingMessageState::Running;
1735 }
1736
1737 fn is_empty(&self) -> bool {
1738 self.queued.is_empty()
1739 }
1740}
1741
1742#[derive(Inspect)]
1743struct ClientTaskInner {
1744 messages: OutgoingMessages,
1745 #[inspect(with = "|x| inspect::iter_by_key(x).map_key(|id| id.0)")]
1746 teardown_gpadls: HashMap<GpadlId, ChannelId>,
1747 #[inspect(skip)]
1748 channel_requests: SelectAll<TaggedStream<ChannelId, mesh::Receiver<ChannelRequest>>>,
1749 synic: SynicState,
1750}
1751
1752#[derive(Inspect)]
1753struct SynicState {
1754 #[inspect(skip)]
1755 event_client: Arc<dyn SynicEventClient>,
1756 #[inspect(iter_by_index)]
1757 event_flag_state: Vec<bool>,
1758}
1759
1760#[derive(Inspect, Default)]
1761#[inspect(transparent)]
1762struct ChannelList(
1763 #[inspect(with = "|x| inspect::iter_by_key(x).map_key(|id| id.0)")] HashMap<ChannelId, Channel>,
1764);
1765
1766struct ChannelRef<'a>(hash_map::OccupiedEntry<'a, ChannelId, Channel>);
1769
1770struct TriedRelease(());
1774
1775impl ChannelRef<'_> {
1776 fn try_release(self, messages: &mut OutgoingMessages) -> TriedRelease {
1780 if self.is_client_released
1781 && matches!(self.state, ChannelState::Revoked)
1782 && self.pending_request().is_none()
1783 {
1784 let channel_id = *self.0.key();
1785 tracelimit::info_ratelimited!(
1786 channel_id = channel_id.0,
1787 key = %OfferKey::from(&self.offer),
1788 "releasing channel"
1789 );
1790
1791 messages.send(&protocol::RelIdReleased { channel_id });
1792 self.0.remove();
1793 }
1794 TriedRelease(())
1795 }
1796}
1797
1798impl Deref for ChannelRef<'_> {
1799 type Target = Channel;
1800
1801 fn deref(&self) -> &Self::Target {
1802 self.0.get()
1803 }
1804}
1805
1806impl DerefMut for ChannelRef<'_> {
1807 fn deref_mut(&mut self) -> &mut Self::Target {
1808 self.0.get_mut()
1809 }
1810}
1811
1812impl ChannelList {
1813 fn revoked_channel_with_pending_request(&self) -> Option<(ChannelId, &'static str)> {
1814 self.0.iter().find_map(|(&id, channel)| {
1815 if !matches!(channel.state, ChannelState::Revoked) {
1816 return None;
1817 }
1818 Some((id, channel.pending_request()?))
1819 })
1820 }
1821
1822 #[track_caller]
1823 fn get_mut(&mut self, channel_id: ChannelId) -> ChannelRef<'_> {
1824 match self.0.entry(channel_id) {
1825 hash_map::Entry::Occupied(entry) => ChannelRef(entry),
1826 hash_map::Entry::Vacant(_) => {
1827 panic!("channel {:?} not found", channel_id);
1828 }
1829 }
1830 }
1831}
1832
1833impl SynicState {
1834 fn guest_to_host_interrupt(&self, connection_id: Arc<AtomicU32>) -> Interrupt {
1835 Interrupt::from_fn({
1836 let event_client = self.event_client.clone();
1837 move || {
1838 let connection_id = connection_id.load(Ordering::Acquire);
1839 if connection_id == 0 {
1840 tracing::debug!("interrupt signal after close");
1841 return;
1842 }
1843
1844 if let Err(err) = event_client.signal_event(connection_id, 0) {
1845 tracelimit::warn_ratelimited!(
1846 error = &err as &dyn std::error::Error,
1847 "failed to signal event"
1848 );
1849 }
1850 }
1851 })
1852 }
1853
1854 const MAX_EVENT_FLAGS: u16 = 2047;
1855
1856 fn allocate_event_flag(&mut self, event: &Event) -> Result<u16> {
1857 let i = self
1858 .event_flag_state
1859 .iter()
1860 .position(|&used| !used)
1861 .ok_or(())
1862 .or_else(|()| {
1863 if self.event_flag_state.len() >= Self::MAX_EVENT_FLAGS as usize {
1864 anyhow::bail!("out of event flags");
1865 }
1866 self.event_flag_state.push(false);
1867 Ok(self.event_flag_state.len() - 1)
1868 })?;
1869
1870 let event_flag = (i + 1) as u16;
1871 self.event_client
1872 .map_event(event_flag, event)
1873 .context("failed to map event")?;
1874 self.event_flag_state[i] = true;
1875 Ok(event_flag)
1876 }
1877
1878 fn restore_event_flag(&mut self, flag: u16, event: &Event) -> Result<()> {
1879 let i = (flag as usize)
1880 .checked_sub(1)
1881 .context("invalid event flag")?;
1882 if i >= Self::MAX_EVENT_FLAGS as usize {
1883 anyhow::bail!("invalid event flag");
1884 }
1885 if self.event_flag_state.len() <= i {
1886 self.event_flag_state.resize(i + 1, false);
1887 }
1888 if self.event_flag_state[i] {
1889 anyhow::bail!("event flag already in use");
1890 }
1891 self.event_client
1892 .map_event(flag, event)
1893 .context("failed to map event")?;
1894 self.event_flag_state[i] = true;
1895 Ok(())
1896 }
1897
1898 fn free_event_flag(&mut self, flag: u16) {
1899 let i = flag as usize - 1;
1900 assert!(i < self.event_flag_state.len());
1901 self.event_flag_state[i] = false;
1902 }
1903}
1904
1905#[cfg(test)]
1906mod tests {
1907 use super::*;
1908 use futures_concurrency::future::Join;
1909 use guid::Guid;
1910 use pal_async::DefaultDriver;
1911 use pal_async::async_test;
1912 use pal_async::timer::PolledTimer;
1913 use protocol::TargetInfo;
1914 use std::fmt::Debug;
1915 use std::task::ready;
1916 use std::time::Duration;
1917 use test_with_tracing::test;
1918 use vmbus_core::protocol::MessageHeader;
1919 use vmbus_core::protocol::MessageType;
1920 use vmbus_core::protocol::OfferFlags;
1921 use vmbus_core::protocol::UserDefinedData;
1922 use vmbus_core::protocol::VmbusMessage;
1923 use zerocopy::FromBytes;
1924 use zerocopy::FromZeros;
1925 use zerocopy::Immutable;
1926 use zerocopy::IntoBytes;
1927 use zerocopy::KnownLayout;
1928
1929 const VMBUS_TEST_CLIENT_ID: Guid = guid::guid!("e6e6e6e6-e6e6-e6e6-e6e6-e6e6e6e6e6e6");
1930
1931 fn in_msg<T: IntoBytes + Immutable + KnownLayout>(message_type: MessageType, t: T) -> Vec<u8> {
1932 let mut data = Vec::new();
1933 data.extend_from_slice(&message_type.0.to_ne_bytes());
1934 data.extend_from_slice(&0u32.to_ne_bytes());
1935 data.extend_from_slice(t.as_bytes());
1936 data
1937 }
1938
1939 #[track_caller]
1940 fn check_message<T>(msg: OutgoingMessage, chk: T)
1941 where
1942 T: IntoBytes + FromBytes + Immutable + KnownLayout + Debug + VmbusMessage,
1943 {
1944 check_message_with_data(msg, chk, &[]);
1945 }
1946
1947 #[track_caller]
1948 fn check_message_with_data<T>(msg: OutgoingMessage, chk: T, data: &[u8])
1949 where
1950 T: IntoBytes + FromBytes + Immutable + KnownLayout + Debug + VmbusMessage,
1951 {
1952 let chk_data = OutgoingMessage::with_data(&chk, data);
1953 if msg.data() != chk_data.data() {
1954 let (header, rest) = MessageHeader::read_from_prefix(msg.data()).unwrap();
1955 assert_eq!(header.message_type(), <T as VmbusMessage>::MESSAGE_TYPE);
1956 let (msg, rest) = T::read_from_prefix(rest).expect("incorrect message size");
1957 if msg.as_bytes() != chk.as_bytes() {
1958 panic!("mismatched messages, expected {:#?}, got {:#?}", chk, msg);
1959 }
1960 if rest != data {
1961 panic!("mismatched data, expected {:#?}, got {:#?}", data, rest);
1962 }
1963 }
1964 }
1965
1966 struct TestServer {
1967 messages: mesh::Receiver<OutgoingMessage>,
1968 send: mesh::Sender<Vec<u8>>,
1969 }
1970
1971 impl TestServer {
1972 async fn next(&mut self) -> Option<OutgoingMessage> {
1973 self.messages.next().await
1974 }
1975
1976 fn send(&self, msg: Vec<u8>) {
1977 self.send.send(msg);
1978 }
1979
1980 async fn connect(&mut self, client: &mut VmbusClient) -> ConnectResult {
1981 self.connect_with_channels(client, |_| {}).await
1982 }
1983
1984 async fn connect_with_channels(
1985 &mut self,
1986 client: &mut VmbusClient,
1987 send_offers: impl FnOnce(&mut Self),
1988 ) -> ConnectResult {
1989 let client_connect = client.connect(0, None, Guid::ZERO);
1990
1991 let server_connect = async {
1992 let _ = self.next().await.unwrap();
1993
1994 self.send(in_msg(
1995 MessageType::VERSION_RESPONSE,
1996 protocol::VersionResponse2 {
1997 version_response: protocol::VersionResponse {
1998 version_supported: 1,
1999 connection_state: ConnectionState::SUCCESSFUL,
2000 padding: 0,
2001 selected_version_or_connection_id: 0,
2002 },
2003 supported_features: SUPPORTED_FEATURE_FLAGS.into(),
2004 },
2005 ));
2006
2007 check_message(self.next().await.unwrap(), protocol::RequestOffers {});
2008
2009 send_offers(self);
2010 self.send(in_msg(MessageType::ALL_OFFERS_DELIVERED, [0x00]));
2011 };
2012
2013 let (connection, ()) = (client_connect, server_connect).join().await;
2014
2015 let connection = connection.unwrap();
2016 assert_eq!(connection.version.version, Version::Copper);
2017 assert_eq!(connection.version.feature_flags, SUPPORTED_FEATURE_FLAGS);
2018 connection
2019 }
2020
2021 async fn get_channel(&mut self, client: &mut VmbusClient) -> OfferInfo {
2022 let [channel] = self
2023 .get_channels(client, 1)
2024 .await
2025 .offers
2026 .try_into()
2027 .unwrap();
2028 channel
2029 }
2030
2031 async fn get_channels(&mut self, client: &mut VmbusClient, count: usize) -> ConnectResult {
2032 self.connect_with_channels(client, |this| {
2033 for i in 0..count {
2034 let offer = protocol::OfferChannel {
2035 interface_id: Guid::new_random(),
2036 instance_id: Guid::new_random(),
2037 rsvd: [0; 4],
2038 flags: OfferFlags::new(),
2039 mmio_megabytes: 0,
2040 user_defined: UserDefinedData::new_zeroed(),
2041 subchannel_index: 0,
2042 mmio_megabytes_optional: 0,
2043 channel_id: ChannelId(i as u32),
2044 monitor_id: 0,
2045 monitor_allocated: 0,
2046 is_dedicated: 0,
2047 connection_id: 0,
2048 };
2049
2050 this.send(in_msg(MessageType::OFFER_CHANNEL, offer));
2051 }
2052 })
2053 .await
2054 }
2055
2056 async fn stop_client(&mut self, client: &mut VmbusClient) {
2057 let client_stop = client.stop();
2058 let server_stop = async {
2059 check_message(self.next().await.unwrap(), protocol::Pause);
2060 self.send(in_msg(MessageType::PAUSE_RESPONSE, protocol::PauseResponse));
2061 };
2062 (client_stop, server_stop).join().await;
2063 }
2064
2065 async fn start_client(&mut self, client: &mut VmbusClient) {
2066 client.start();
2067 check_message(self.next().await.unwrap(), protocol::Resume);
2068 }
2069 }
2070
2071 struct TestServerClient {
2072 sender: mesh::Sender<OutgoingMessage>,
2073 timer: PolledTimer,
2074 deadline: Option<pal_async::timer::Instant>,
2075 }
2076
2077 impl PollPostMessage for TestServerClient {
2078 fn poll_post_message(
2079 &mut self,
2080 cx: &mut Context<'_>,
2081 _connection_id: u32,
2082 _typ: u32,
2083 msg: &[u8],
2084 ) -> Poll<()> {
2085 loop {
2086 if let Some(deadline) = self.deadline {
2087 ready!(self.timer.poll_until(cx, deadline));
2088 self.deadline = None;
2089 }
2090 let mut b = [0];
2095 getrandom::fill(&mut b).unwrap();
2096 if b[0] % 4 == 0 {
2097 self.deadline =
2098 Some(pal_async::timer::Instant::now() + Duration::from_millis(10));
2099 } else {
2100 let msg = OutgoingMessage::from_message(msg).unwrap();
2101 tracing::info!(
2102 msg = ?MessageHeader::read_from_prefix(msg.data()),
2103 "sending message"
2104 );
2105 self.sender.send(msg);
2106 break Poll::Ready(());
2107 }
2108 }
2109 }
2110 }
2111
2112 struct NoopSynicEvents;
2113
2114 impl SynicEventClient for NoopSynicEvents {
2115 fn map_event(&self, _event_flag: u16, _event: &Event) -> std::io::Result<()> {
2116 Ok(())
2117 }
2118
2119 fn unmap_event(&self, _event_flag: u16) {}
2120
2121 fn signal_event(&self, _connection_id: u32, _event_flag: u16) -> std::io::Result<()> {
2122 Err(std::io::ErrorKind::Unsupported.into())
2123 }
2124 }
2125
2126 struct TestMessageSource {
2127 msg_recv: mesh::Receiver<Vec<u8>>,
2128 paused: bool,
2129 }
2130
2131 impl AsyncRecv for TestMessageSource {
2132 fn poll_recv(
2133 &mut self,
2134 cx: &mut Context<'_>,
2135 mut bufs: &mut [std::io::IoSliceMut<'_>],
2136 ) -> Poll<std::io::Result<usize>> {
2137 let value = match self.msg_recv.poll_recv(cx) {
2138 Poll::Ready(v) => v.unwrap(),
2139 Poll::Pending => {
2140 if self.paused {
2141 return Poll::Ready(Ok(0));
2142 } else {
2143 return Poll::Pending;
2144 }
2145 }
2146 };
2147 let mut remaining = value.as_slice();
2148 let mut total_size = 0;
2149 while !remaining.is_empty() && !bufs.is_empty() {
2150 let size = bufs[0].len().min(remaining.len());
2151 bufs[0][..size].copy_from_slice(&remaining[..size]);
2152 remaining = &remaining[size..];
2153 bufs = &mut bufs[1..];
2154 total_size += size;
2155 }
2156
2157 Ok(total_size).into()
2158 }
2159 }
2160
2161 impl VmbusMessageSource for TestMessageSource {
2162 fn pause_message_stream(&mut self) {
2163 self.paused = true;
2164 }
2165
2166 fn resume_message_stream(&mut self) {
2167 self.paused = false;
2168 }
2169 }
2170
2171 fn test_init(driver: &DefaultDriver) -> (TestServer, VmbusClient) {
2172 let (msg_send, msg_recv) = mesh::channel();
2173 let (synic_send, synic_recv) = mesh::channel();
2174 let server = TestServer {
2175 messages: synic_recv,
2176 send: msg_send,
2177 };
2178 let mut client = VmbusClientBuilder::new(
2179 NoopSynicEvents,
2180 TestMessageSource {
2181 msg_recv,
2182 paused: false,
2183 },
2184 TestServerClient {
2185 sender: synic_send,
2186 deadline: None,
2187 timer: PolledTimer::new(driver),
2188 },
2189 )
2190 .build(driver);
2191 client.start();
2192 (server, client)
2193 }
2194
2195 #[async_test]
2196 async fn test_initiate_contact_success(driver: DefaultDriver) {
2197 let (mut server, client) = test_init(&driver);
2198 let _recv = client
2199 .access
2200 .client_request_send
2201 .call(ClientRequest::Connect, ConnectRequest::default());
2202 check_message(
2203 server.next().await.unwrap(),
2204 protocol::InitiateContact2 {
2205 initiate_contact: protocol::InitiateContact {
2206 version_requested: Version::Copper as u32,
2207 target_message_vp: 0,
2208 interrupt_page_or_target_info: TargetInfo::new()
2209 .with_sint(2)
2210 .with_vtl(0)
2211 .with_feature_flags(SUPPORTED_FEATURE_FLAGS.into())
2212 .into(),
2213 parent_to_child_monitor_page_gpa: 0,
2214 child_to_parent_monitor_page_gpa: 0,
2215 },
2216 ..FromZeros::new_zeroed()
2217 },
2218 );
2219 }
2220
2221 #[async_test]
2222 async fn test_connect_success(driver: DefaultDriver) {
2223 let (mut server, mut client) = test_init(&driver);
2224 let client_connect = client.connect(0, None, Guid::ZERO);
2225
2226 let server_connect = async {
2227 check_message(
2228 server.next().await.unwrap(),
2229 protocol::InitiateContact2 {
2230 initiate_contact: protocol::InitiateContact {
2231 version_requested: Version::Copper as u32,
2232 target_message_vp: 0,
2233 interrupt_page_or_target_info: TargetInfo::new()
2234 .with_sint(2)
2235 .with_vtl(0)
2236 .with_feature_flags(SUPPORTED_FEATURE_FLAGS.into())
2237 .into(),
2238 parent_to_child_monitor_page_gpa: 0,
2239 child_to_parent_monitor_page_gpa: 0,
2240 },
2241 ..FromZeros::new_zeroed()
2242 },
2243 );
2244
2245 server.send(in_msg(
2246 MessageType::VERSION_RESPONSE,
2247 protocol::VersionResponse2 {
2248 version_response: protocol::VersionResponse {
2249 version_supported: 1,
2250 connection_state: ConnectionState::SUCCESSFUL,
2251 padding: 0,
2252 selected_version_or_connection_id: 0,
2253 },
2254 supported_features: SUPPORTED_FEATURE_FLAGS.into_bits(),
2255 },
2256 ));
2257
2258 check_message(server.next().await.unwrap(), protocol::RequestOffers {});
2259 server.send(in_msg(MessageType::ALL_OFFERS_DELIVERED, [0x00]));
2260 };
2261
2262 let (connection, ()) = (client_connect, server_connect).join().await;
2263 let connection = connection.unwrap();
2264
2265 assert_eq!(connection.version.version, Version::Copper);
2266 assert_eq!(connection.version.feature_flags, SUPPORTED_FEATURE_FLAGS);
2267 }
2268
2269 #[async_test]
2270 async fn test_feature_flags(driver: DefaultDriver) {
2271 let (mut server, mut client) = test_init(&driver);
2272 let client_connect = client.connect(0, None, Guid::ZERO);
2273
2274 let server_connect = async {
2275 check_message(
2276 server.next().await.unwrap(),
2277 protocol::InitiateContact2 {
2278 initiate_contact: protocol::InitiateContact {
2279 version_requested: Version::Copper as u32,
2280 target_message_vp: 0,
2281 interrupt_page_or_target_info: TargetInfo::new()
2282 .with_sint(2)
2283 .with_vtl(0)
2284 .with_feature_flags(SUPPORTED_FEATURE_FLAGS.into())
2285 .into(),
2286 parent_to_child_monitor_page_gpa: 0,
2287 child_to_parent_monitor_page_gpa: 0,
2288 },
2289 ..FromZeros::new_zeroed()
2290 },
2291 );
2292
2293 server.send(in_msg(
2296 MessageType::VERSION_RESPONSE,
2297 protocol::VersionResponse2 {
2298 version_response: protocol::VersionResponse {
2299 version_supported: 1,
2300 connection_state: ConnectionState::SUCCESSFUL,
2301 padding: 0,
2302 selected_version_or_connection_id: 0,
2303 },
2304 supported_features: 2,
2305 },
2306 ));
2307
2308 check_message(server.next().await.unwrap(), protocol::RequestOffers {});
2309 server.send(in_msg(MessageType::ALL_OFFERS_DELIVERED, [0x00]));
2310 };
2311
2312 let (connection, ()) = (client_connect, server_connect).join().await;
2313 let connection = connection.unwrap();
2314
2315 assert_eq!(connection.version.version, Version::Copper);
2316 assert_eq!(
2317 connection.version.feature_flags,
2318 FeatureFlags::new().with_channel_interrupt_redirection(true)
2319 );
2320 }
2321
2322 #[async_test]
2323 async fn test_client_id(driver: DefaultDriver) {
2324 let (mut server, client) = test_init(&driver);
2325 let initiate_contact = ConnectRequest {
2326 client_id: VMBUS_TEST_CLIENT_ID,
2327 ..Default::default()
2328 };
2329 let _recv = client
2330 .access
2331 .client_request_send
2332 .call(ClientRequest::Connect, initiate_contact);
2333
2334 check_message(
2335 server.next().await.unwrap(),
2336 protocol::InitiateContact2 {
2337 initiate_contact: protocol::InitiateContact {
2338 version_requested: Version::Copper as u32,
2339 target_message_vp: 0,
2340 interrupt_page_or_target_info: TargetInfo::new()
2341 .with_sint(2)
2342 .with_vtl(0)
2343 .with_feature_flags(SUPPORTED_FEATURE_FLAGS.into())
2344 .into(),
2345 parent_to_child_monitor_page_gpa: 0,
2346 child_to_parent_monitor_page_gpa: 0,
2347 },
2348 client_id: VMBUS_TEST_CLIENT_ID,
2349 },
2350 );
2351 }
2352
2353 #[async_test]
2354 async fn test_version_negotiation(driver: DefaultDriver) {
2355 let (mut server, mut client) = test_init(&driver);
2356 let client_connect = client.connect(0, None, Guid::ZERO);
2357
2358 let server_connect = async {
2359 check_message(
2360 server.next().await.unwrap(),
2361 protocol::InitiateContact2 {
2362 initiate_contact: protocol::InitiateContact {
2363 version_requested: Version::Copper as u32,
2364 target_message_vp: 0,
2365 interrupt_page_or_target_info: TargetInfo::new()
2366 .with_sint(2)
2367 .with_vtl(0)
2368 .with_feature_flags(SUPPORTED_FEATURE_FLAGS.into())
2369 .into(),
2370 parent_to_child_monitor_page_gpa: 0,
2371 child_to_parent_monitor_page_gpa: 0,
2372 },
2373 ..FromZeros::new_zeroed()
2374 },
2375 );
2376
2377 server.send(in_msg(
2378 MessageType::VERSION_RESPONSE,
2379 protocol::VersionResponse {
2380 version_supported: 0,
2381 connection_state: ConnectionState::SUCCESSFUL,
2382 padding: 0,
2383 selected_version_or_connection_id: 0,
2384 },
2385 ));
2386
2387 check_message(
2388 server.next().await.unwrap(),
2389 protocol::InitiateContact {
2390 version_requested: Version::Iron as u32,
2391 target_message_vp: 0,
2392 interrupt_page_or_target_info: TargetInfo::new()
2393 .with_sint(2)
2394 .with_vtl(0)
2395 .with_feature_flags(FeatureFlags::new().into())
2396 .into(),
2397 parent_to_child_monitor_page_gpa: 0,
2398 child_to_parent_monitor_page_gpa: 0,
2399 },
2400 );
2401
2402 server.send(in_msg(
2403 MessageType::VERSION_RESPONSE,
2404 protocol::VersionResponse {
2405 version_supported: 1,
2406 connection_state: ConnectionState::SUCCESSFUL,
2407 padding: 0,
2408 selected_version_or_connection_id: 0,
2409 },
2410 ));
2411
2412 check_message(server.next().await.unwrap(), protocol::RequestOffers {});
2413 server.send(in_msg(MessageType::ALL_OFFERS_DELIVERED, [0x00]));
2414 };
2415
2416 let (connection, ()) = (client_connect, server_connect).join().await;
2417 let connection = connection.unwrap();
2418
2419 assert_eq!(connection.version.version, Version::Iron);
2420 assert_eq!(connection.version.feature_flags, FeatureFlags::new());
2421 }
2422
2423 #[async_test]
2424 async fn test_open_channel_success(driver: DefaultDriver) {
2425 let (mut server, mut client) = test_init(&driver);
2426 let channel = server.get_channel(&mut client).await;
2427
2428 let recv = channel.request_send.call(
2429 ChannelRequest::Open,
2430 OpenRequest {
2431 open_data: OpenData {
2432 target_vp: Some(0),
2433 ring_offset: 0,
2434 ring_gpadl_id: GpadlId(0),
2435 event_flag: 0,
2436 connection_id: 0,
2437 user_data: UserDefinedData::new_zeroed(),
2438 },
2439 incoming_event: None,
2440 use_vtl2_connection_id: false,
2441 },
2442 );
2443
2444 check_message(
2445 server.next().await.unwrap(),
2446 protocol::OpenChannel2 {
2447 open_channel: protocol::OpenChannel {
2448 channel_id: ChannelId(0),
2449 open_id: 0,
2450 ring_buffer_gpadl_id: GpadlId(0),
2451 target_vp: 0,
2452 downstream_ring_buffer_page_offset: 0,
2453 user_data: UserDefinedData::new_zeroed(),
2454 },
2455 connection_id: 0,
2456 event_flag: 0,
2457 flags: Default::default(),
2458 },
2459 );
2460
2461 server.send(in_msg(
2462 MessageType::OPEN_CHANNEL_RESULT,
2463 protocol::OpenResult {
2464 channel_id: ChannelId(0),
2465 open_id: 0,
2466 status: protocol::STATUS_SUCCESS as u32,
2467 },
2468 ));
2469
2470 recv.await.unwrap().unwrap();
2471 }
2472
2473 #[async_test]
2474 async fn test_open_channel_fail(driver: DefaultDriver) {
2475 let (mut server, mut client) = test_init(&driver);
2476 let channel = server.get_channel(&mut client).await;
2477
2478 let recv = channel.request_send.call(
2479 ChannelRequest::Open,
2480 OpenRequest {
2481 open_data: OpenData {
2482 target_vp: Some(0),
2483 ring_offset: 0,
2484 ring_gpadl_id: GpadlId(0),
2485 event_flag: 0,
2486 connection_id: 0,
2487 user_data: UserDefinedData::new_zeroed(),
2488 },
2489 incoming_event: None,
2490 use_vtl2_connection_id: false,
2491 },
2492 );
2493
2494 check_message(
2495 server.next().await.unwrap(),
2496 protocol::OpenChannel2 {
2497 open_channel: protocol::OpenChannel {
2498 channel_id: ChannelId(0),
2499 open_id: 0,
2500 ring_buffer_gpadl_id: GpadlId(0),
2501 target_vp: 0,
2502 downstream_ring_buffer_page_offset: 0,
2503 user_data: UserDefinedData::new_zeroed(),
2504 },
2505 connection_id: 0,
2506 event_flag: 0,
2507 flags: Default::default(),
2508 },
2509 );
2510
2511 server.send(in_msg(
2512 MessageType::OPEN_CHANNEL_RESULT,
2513 protocol::OpenResult {
2514 channel_id: ChannelId(0),
2515 open_id: 0,
2516 status: protocol::STATUS_UNSUCCESSFUL as u32,
2517 },
2518 ));
2519
2520 recv.await.unwrap().unwrap_err();
2521 }
2522
2523 #[async_test]
2524 async fn test_modify_channel(driver: DefaultDriver) {
2525 let (mut server, mut client) = test_init(&driver);
2526 let channel = server.get_channel(&mut client).await;
2527
2528 let recv = channel.request_send.call(
2531 ChannelRequest::Modify,
2532 ModifyRequest::TargetVp { target_vp: 1 },
2533 );
2534
2535 check_message(
2536 server.next().await.unwrap(),
2537 protocol::ModifyChannel {
2538 channel_id: ChannelId(0),
2539 target_vp: 1,
2540 },
2541 );
2542
2543 server.send(in_msg(
2544 MessageType::MODIFY_CHANNEL_RESPONSE,
2545 protocol::ModifyChannelResponse {
2546 channel_id: ChannelId(0),
2547 status: protocol::STATUS_SUCCESS,
2548 },
2549 ));
2550
2551 let status = recv.await.unwrap();
2552 assert_eq!(status, protocol::STATUS_SUCCESS);
2553 }
2554
2555 #[async_test]
2556 async fn test_save_restore_connected(driver: DefaultDriver) {
2557 let (mut server, mut client) = test_init(&driver);
2558 server.connect(&mut client).await;
2559 server.stop_client(&mut client).await;
2560 let s0 = client.save().await;
2561 let builder = client.sever().await;
2562 let mut client = builder.build(&driver);
2563 client.restore(s0.clone()).await.unwrap();
2564
2565 let s1 = client.save().await;
2566
2567 assert_eq!(s0, s1);
2568 }
2569
2570 #[async_test]
2571 async fn test_save_restore_connected_with_channel(driver: DefaultDriver) {
2572 let (mut server, mut client) = test_init(&driver);
2573 let c0 = server.get_channel(&mut client).await;
2574 server.stop_client(&mut client).await;
2575 let s0 = client.save().await;
2576 let builder = client.sever().await;
2577 let mut client = builder.build(&driver);
2578 let connection = client.restore(s0.clone()).await.unwrap().unwrap();
2579 let s1 = client.save().await;
2580 assert_eq!(s0, s1);
2581 assert_eq!(connection.offers[0].offer, c0.offer);
2582 }
2583
2584 #[async_test]
2585 async fn test_save_restore_connected_with_revoked_channel(driver: DefaultDriver) {
2586 let (mut server, mut client) = test_init(&driver);
2587 let c0 = server.get_channel(&mut client).await;
2588 server.send(in_msg(
2589 MessageType::RESCIND_CHANNEL_OFFER,
2590 protocol::RescindChannelOffer {
2591 channel_id: ChannelId(0),
2592 },
2593 ));
2594 c0.revoke_recv.await.unwrap();
2595 let rpc = c0.request_send.call(
2596 ChannelRequest::Modify,
2597 ModifyRequest::TargetVp { target_vp: 1 },
2598 );
2599
2600 check_message(
2601 server.next().await.unwrap(),
2602 protocol::ModifyChannel {
2603 channel_id: ChannelId(0),
2604 target_vp: 1,
2605 },
2606 );
2607
2608 let client_stop = client.stop();
2609 let server_stop = async {
2610 server.send(in_msg(
2611 MessageType::MODIFY_CHANNEL_RESPONSE,
2612 protocol::ModifyChannelResponse {
2613 channel_id: ChannelId(0),
2614 status: protocol::STATUS_SUCCESS,
2615 },
2616 ));
2617 check_message(server.next().await.unwrap(), protocol::Pause);
2618 server.send(in_msg(MessageType::PAUSE_RESPONSE, protocol::PauseResponse));
2619 };
2620 (client_stop, server_stop).join().await;
2621
2622 rpc.await.unwrap();
2623
2624 let s0 = client.save().await;
2625 let builder = client.sever().await;
2626 let mut client = builder.build(&driver);
2627 let connection = client.restore(s0.clone()).await.unwrap().unwrap();
2628 let s1 = client.save().await;
2629 assert_eq!(s0, s1);
2630 assert!(connection.offers.is_empty());
2631 server.start_client(&mut client).await;
2632 check_message(
2633 server.next().await.unwrap(),
2634 protocol::RelIdReleased {
2635 channel_id: ChannelId(0),
2636 },
2637 );
2638 }
2639
2640 #[async_test]
2641 async fn test_connect_fails_on_incorrect_state(driver: DefaultDriver) {
2642 let (mut server, mut client) = test_init(&driver);
2643 server.connect(&mut client).await;
2644 let err = client.connect(0, None, Guid::ZERO).await.unwrap_err();
2645 assert!(matches!(err, ConnectError::InvalidState), "{:?}", err);
2646 }
2647
2648 #[async_test]
2649 async fn test_hot_add_remove(driver: DefaultDriver) {
2650 let (mut server, mut client) = test_init(&driver);
2651
2652 let mut connection = server.connect(&mut client).await;
2653 let offer = protocol::OfferChannel {
2654 interface_id: Guid::new_random(),
2655 instance_id: Guid::new_random(),
2656 rsvd: [0; 4],
2657 flags: OfferFlags::new(),
2658 mmio_megabytes: 0,
2659 user_defined: UserDefinedData::new_zeroed(),
2660 subchannel_index: 0,
2661 mmio_megabytes_optional: 0,
2662 channel_id: ChannelId(5),
2663 monitor_id: 0,
2664 monitor_allocated: 0,
2665 is_dedicated: 0,
2666 connection_id: 0,
2667 };
2668
2669 server.send(in_msg(MessageType::OFFER_CHANNEL, offer));
2670 let info = connection.offer_recv.next().await.unwrap();
2671
2672 assert_eq!(offer, info.offer);
2673
2674 server.send(in_msg(
2675 MessageType::RESCIND_CHANNEL_OFFER,
2676 protocol::RescindChannelOffer {
2677 channel_id: ChannelId(5),
2678 },
2679 ));
2680
2681 info.revoke_recv.await.unwrap();
2682 drop(info.request_send);
2683
2684 check_message(
2685 server.next().await.unwrap(),
2686 protocol::RelIdReleased {
2687 channel_id: ChannelId(5),
2688 },
2689 );
2690 }
2691
2692 #[async_test]
2693 async fn test_gpadl_success(driver: DefaultDriver) {
2694 let (mut server, mut client) = test_init(&driver);
2695 let channel = server.get_channel(&mut client).await;
2696 let recv = channel.request_send.call(
2697 ChannelRequest::Gpadl,
2698 GpadlRequest {
2699 id: GpadlId(1),
2700 count: 1,
2701 buf: vec![5],
2702 },
2703 );
2704
2705 check_message_with_data(
2706 server.next().await.unwrap(),
2707 protocol::GpadlHeader {
2708 channel_id: ChannelId(0),
2709 gpadl_id: GpadlId(1),
2710 len: 8,
2711 count: 1,
2712 },
2713 0x5u64.as_bytes(),
2714 );
2715
2716 server.send(in_msg(
2717 MessageType::GPADL_CREATED,
2718 protocol::GpadlCreated {
2719 channel_id: ChannelId(0),
2720 gpadl_id: GpadlId(1),
2721 status: protocol::STATUS_SUCCESS,
2722 },
2723 ));
2724
2725 recv.await.unwrap().unwrap();
2726
2727 let rpc = channel
2728 .request_send
2729 .call(ChannelRequest::TeardownGpadl, GpadlId(1));
2730
2731 check_message(
2732 server.next().await.unwrap(),
2733 protocol::GpadlTeardown {
2734 channel_id: ChannelId(0),
2735 gpadl_id: GpadlId(1),
2736 },
2737 );
2738
2739 server.send(in_msg(
2740 MessageType::GPADL_TORNDOWN,
2741 protocol::GpadlTorndown {
2742 gpadl_id: GpadlId(1),
2743 },
2744 ));
2745
2746 rpc.await.unwrap();
2747 }
2748
2749 #[async_test]
2750 async fn test_gpadl_fail(driver: DefaultDriver) {
2751 let (mut server, mut client) = test_init(&driver);
2752 let channel = server.get_channel(&mut client).await;
2753 let recv = channel.request_send.call(
2754 ChannelRequest::Gpadl,
2755 GpadlRequest {
2756 id: GpadlId(1),
2757 count: 1,
2758 buf: vec![7],
2759 },
2760 );
2761
2762 check_message_with_data(
2763 server.next().await.unwrap(),
2764 protocol::GpadlHeader {
2765 channel_id: ChannelId(0),
2766 gpadl_id: GpadlId(1),
2767 len: 8,
2768 count: 1,
2769 },
2770 0x7u64.as_bytes(),
2771 );
2772
2773 server.send(in_msg(
2774 MessageType::GPADL_CREATED,
2775 protocol::GpadlCreated {
2776 channel_id: ChannelId(0),
2777 gpadl_id: GpadlId(1),
2778 status: protocol::STATUS_UNSUCCESSFUL,
2779 },
2780 ));
2781
2782 recv.await.unwrap().unwrap_err();
2783 }
2784
2785 #[async_test]
2786 async fn test_gpadl_with_revoke(driver: DefaultDriver) {
2787 let (mut server, mut client) = test_init(&driver);
2788 let channel = server.get_channel(&mut client).await;
2789 let channel_id = ChannelId(0);
2790 for gpadl_id in [1, 2, 3].map(GpadlId) {
2791 let recv = channel.request_send.call(
2792 ChannelRequest::Gpadl,
2793 GpadlRequest {
2794 id: gpadl_id,
2795 count: 1,
2796 buf: vec![3],
2797 },
2798 );
2799
2800 check_message_with_data(
2801 server.next().await.unwrap(),
2802 protocol::GpadlHeader {
2803 channel_id,
2804 gpadl_id,
2805 len: 8,
2806 count: 1,
2807 },
2808 0x3u64.as_bytes(),
2809 );
2810
2811 server.send(in_msg(
2812 MessageType::GPADL_CREATED,
2813 protocol::GpadlCreated {
2814 channel_id,
2815 gpadl_id,
2816 status: protocol::STATUS_SUCCESS,
2817 },
2818 ));
2819
2820 recv.await.unwrap().unwrap();
2821 }
2822
2823 let rpc = channel
2824 .request_send
2825 .call(ChannelRequest::TeardownGpadl, GpadlId(1));
2826
2827 check_message(
2828 server.next().await.unwrap(),
2829 protocol::GpadlTeardown {
2830 channel_id,
2831 gpadl_id: GpadlId(1),
2832 },
2833 );
2834
2835 server.send(in_msg(
2836 MessageType::RESCIND_CHANNEL_OFFER,
2837 protocol::RescindChannelOffer { channel_id },
2838 ));
2839
2840 let recv = channel.request_send.call_failable(
2841 ChannelRequest::Gpadl,
2842 GpadlRequest {
2843 id: GpadlId(4),
2844 count: 1,
2845 buf: vec![3],
2846 },
2847 );
2848
2849 check_message_with_data(
2850 server.next().await.unwrap(),
2851 protocol::GpadlHeader {
2852 channel_id,
2853 gpadl_id: GpadlId(4),
2854 len: 8,
2855 count: 1,
2856 },
2857 0x3u64.as_bytes(),
2858 );
2859
2860 server.send(in_msg(
2861 MessageType::GPADL_CREATED,
2862 protocol::GpadlCreated {
2863 channel_id,
2864 gpadl_id: GpadlId(4),
2865 status: protocol::STATUS_UNSUCCESSFUL,
2866 },
2867 ));
2868
2869 server.send(in_msg(
2870 MessageType::GPADL_TORNDOWN,
2871 protocol::GpadlTorndown {
2872 gpadl_id: GpadlId(1),
2873 },
2874 ));
2875
2876 rpc.await.unwrap();
2877 recv.await.unwrap_err();
2878
2879 channel.revoke_recv.await.unwrap();
2880
2881 let rpc = channel
2882 .request_send
2883 .call(ChannelRequest::TeardownGpadl, GpadlId(2));
2884 drop(channel.request_send);
2885
2886 check_message(
2887 server.next().await.unwrap(),
2888 protocol::GpadlTeardown {
2889 channel_id,
2890 gpadl_id: GpadlId(2),
2891 },
2892 );
2893
2894 server.send(in_msg(
2895 MessageType::GPADL_TORNDOWN,
2896 protocol::GpadlTorndown {
2897 gpadl_id: GpadlId(2),
2898 },
2899 ));
2900
2901 rpc.await.unwrap();
2902
2903 check_message(
2904 server.next().await.unwrap(),
2905 protocol::RelIdReleased { channel_id },
2906 );
2907 }
2908
2909 #[async_test]
2910 async fn test_modify_connection(driver: DefaultDriver) {
2911 let (mut server, mut client) = test_init(&driver);
2912 server.connect(&mut client).await;
2913 let call = client.access.client_request_send.call(
2914 ClientRequest::Modify,
2915 ModifyConnectionRequest {
2916 monitor_page: Some(MonitorPageGpas {
2917 child_to_parent: 5,
2918 parent_to_child: 6,
2919 }),
2920 },
2921 );
2922
2923 check_message(
2924 server.next().await.unwrap(),
2925 protocol::ModifyConnection {
2926 child_to_parent_monitor_page_gpa: 5,
2927 parent_to_child_monitor_page_gpa: 6,
2928 },
2929 );
2930
2931 server.send(in_msg(
2932 MessageType::MODIFY_CONNECTION_RESPONSE,
2933 protocol::ModifyConnectionResponse {
2934 connection_state: ConnectionState::FAILED_LOW_RESOURCES,
2935 },
2936 ));
2937
2938 let result = call.await.unwrap();
2939 assert_eq!(ConnectionState::FAILED_LOW_RESOURCES, result);
2940 }
2941
2942 #[async_test]
2943 async fn test_hvsock(driver: DefaultDriver) {
2944 let (mut server, mut client) = test_init(&driver);
2945 server.connect(&mut client).await;
2946 let request = HvsockConnectRequest {
2947 service_id: Guid::new_random(),
2948 endpoint_id: Guid::new_random(),
2949 silo_id: Guid::new_random(),
2950 hosted_silo_unaware: false,
2951 };
2952
2953 let resp = client.access().connect_hvsock(request);
2954 check_message(
2955 server.next().await.unwrap(),
2956 protocol::TlConnectRequest2 {
2957 base: protocol::TlConnectRequest {
2958 service_id: request.service_id,
2959 endpoint_id: request.endpoint_id,
2960 },
2961 silo_id: request.silo_id,
2962 },
2963 );
2964
2965 server.send(in_msg(
2967 MessageType::TL_CONNECT_REQUEST_RESULT,
2968 protocol::TlConnectResult {
2969 service_id: request.service_id,
2970 endpoint_id: request.endpoint_id,
2971 status: protocol::STATUS_CONNECTION_REFUSED,
2972 },
2973 ));
2974
2975 let result = resp.await;
2976 assert!(result.is_none());
2977 }
2978
2979 #[async_test]
2980 async fn test_synic_event_flags(driver: DefaultDriver) {
2981 let (mut server, mut client) = test_init(&driver);
2982 let connection = server.get_channels(&mut client, 5).await;
2983 let event = Event::new();
2984
2985 for _ in 0..5 {
2986 for (i, channel) in connection.offers.iter().enumerate() {
2987 let recv = channel.request_send.call(
2988 ChannelRequest::Open,
2989 OpenRequest {
2990 open_data: OpenData {
2991 target_vp: Some(0),
2992 ring_offset: 0,
2993 ring_gpadl_id: GpadlId(0),
2994 event_flag: 0,
2995 connection_id: 0,
2996 user_data: UserDefinedData::new_zeroed(),
2997 },
2998 incoming_event: Some(event.clone()),
2999 use_vtl2_connection_id: false,
3000 },
3001 );
3002
3003 let expected_event_flag = i as u16 + 1;
3004
3005 check_message(
3006 server.next().await.unwrap(),
3007 protocol::OpenChannel2 {
3008 open_channel: protocol::OpenChannel {
3009 channel_id: channel.offer.channel_id,
3010 open_id: 0,
3011 ring_buffer_gpadl_id: GpadlId(0),
3012 target_vp: 0,
3013 downstream_ring_buffer_page_offset: 0,
3014 user_data: UserDefinedData::new_zeroed(),
3015 },
3016 connection_id: 0,
3017 event_flag: expected_event_flag,
3018 flags: OpenChannelFlags::new().with_redirect_interrupt(true),
3019 },
3020 );
3021
3022 server.send(in_msg(
3023 MessageType::OPEN_CHANNEL_RESULT,
3024 protocol::OpenResult {
3025 channel_id: channel.offer.channel_id,
3026 open_id: 0,
3027 status: protocol::STATUS_SUCCESS as u32,
3028 },
3029 ));
3030
3031 let output = recv.await.unwrap().unwrap();
3032 assert_eq!(output.redirected_event_flag, Some(expected_event_flag));
3033 }
3034
3035 for (i, channel) in connection.offers.iter().enumerate() {
3036 channel
3039 .request_send
3040 .call(ChannelRequest::Close, ())
3041 .await
3042 .unwrap();
3043
3044 check_message(
3045 server.next().await.unwrap(),
3046 protocol::CloseChannel {
3047 channel_id: ChannelId(i as u32),
3048 },
3049 );
3050 }
3051 }
3052 }
3053
3054 #[async_test]
3055 async fn test_revoke(driver: DefaultDriver) {
3056 let (mut server, mut client) = test_init(&driver);
3057 let channel = server.get_channel(&mut client).await;
3058
3059 server.send(in_msg(
3060 MessageType::RESCIND_CHANNEL_OFFER,
3061 protocol::RescindChannelOffer {
3062 channel_id: ChannelId(0),
3063 },
3064 ));
3065
3066 channel.revoke_recv.await.unwrap();
3067
3068 channel
3069 .request_send
3070 .call_failable(
3071 ChannelRequest::Open,
3072 OpenRequest {
3073 open_data: OpenData {
3074 target_vp: Some(0),
3075 ring_offset: 0,
3076 ring_gpadl_id: GpadlId(0),
3077 event_flag: 0,
3078 connection_id: 0,
3079 user_data: UserDefinedData::new_zeroed(),
3080 },
3081 incoming_event: None,
3082 use_vtl2_connection_id: false,
3083 },
3084 )
3085 .await
3086 .unwrap_err();
3087 }
3088
3089 #[async_test]
3090 #[should_panic(expected = "channel should not exist")]
3091 async fn test_reoffer_in_use_rel_id(driver: DefaultDriver) {
3092 let (mut server, mut client) = test_init(&driver);
3093 let mut connection = server.get_channels(&mut client, 1).await;
3094 let [channel] = connection.offers.try_into().unwrap();
3095
3096 server.send(in_msg(
3097 MessageType::RESCIND_CHANNEL_OFFER,
3098 protocol::RescindChannelOffer {
3099 channel_id: ChannelId(0),
3100 },
3101 ));
3102
3103 channel.revoke_recv.await.unwrap();
3104
3105 let offer = protocol::OfferChannel {
3107 interface_id: Guid::new_random(),
3108 instance_id: Guid::new_random(),
3109 rsvd: [0; 4],
3110 flags: OfferFlags::new(),
3111 mmio_megabytes: 0,
3112 user_defined: UserDefinedData::new_zeroed(),
3113 subchannel_index: 0,
3114 mmio_megabytes_optional: 0,
3115 channel_id: ChannelId(0),
3116 monitor_id: 0,
3117 monitor_allocated: 0,
3118 is_dedicated: 0,
3119 connection_id: 0,
3120 };
3121
3122 server.send(in_msg(MessageType::OFFER_CHANNEL, offer));
3123
3124 connection.offer_recv.next().await;
3125 }
3126
3127 #[async_test]
3128 async fn test_revoke_release_and_reoffer(driver: DefaultDriver) {
3129 let (mut server, mut client) = test_init(&driver);
3130 let mut connection = server.get_channels(&mut client, 1).await;
3131 let [channel] = connection.offers.try_into().unwrap();
3132
3133 server.send(in_msg(
3134 MessageType::RESCIND_CHANNEL_OFFER,
3135 protocol::RescindChannelOffer {
3136 channel_id: ChannelId(0),
3137 },
3138 ));
3139
3140 channel.revoke_recv.await.unwrap();
3141 drop(channel.request_send);
3142
3143 check_message(
3144 server.next().await.unwrap(),
3145 protocol::RelIdReleased {
3146 channel_id: ChannelId(0),
3147 },
3148 );
3149
3150 let offer = protocol::OfferChannel {
3151 interface_id: Guid::new_random(),
3152 instance_id: Guid::new_random(),
3153 rsvd: [0; 4],
3154 flags: OfferFlags::new(),
3155 mmio_megabytes: 0,
3156 user_defined: UserDefinedData::new_zeroed(),
3157 subchannel_index: 0,
3158 mmio_megabytes_optional: 0,
3159 channel_id: ChannelId(0),
3160 monitor_id: 0,
3161 monitor_allocated: 0,
3162 is_dedicated: 0,
3163 connection_id: 0,
3164 };
3165
3166 server.send(in_msg(MessageType::OFFER_CHANNEL, offer));
3167
3168 connection.offer_recv.next().await.unwrap();
3169 }
3170
3171 #[async_test]
3172 async fn test_release_revoke_and_reoffer(driver: DefaultDriver) {
3173 let (mut server, mut client) = test_init(&driver);
3174 let mut connection = server.get_channels(&mut client, 1).await;
3175 let [channel] = connection.offers.try_into().unwrap();
3176
3177 let open = channel.request_send.call_failable(
3178 ChannelRequest::Open,
3179 OpenRequest {
3180 open_data: OpenData {
3181 target_vp: Some(0),
3182 ring_offset: 0,
3183 ring_gpadl_id: GpadlId(0),
3184 event_flag: 0,
3185 connection_id: 0,
3186 user_data: UserDefinedData::new_zeroed(),
3187 },
3188 incoming_event: None,
3189 use_vtl2_connection_id: false,
3190 },
3191 );
3192
3193 let server_open = async {
3194 check_message(
3195 server.next().await.unwrap(),
3196 protocol::OpenChannel2 {
3197 open_channel: protocol::OpenChannel {
3198 channel_id: ChannelId(0),
3199 open_id: 0,
3200 ring_buffer_gpadl_id: GpadlId(0),
3201 target_vp: 0,
3202 downstream_ring_buffer_page_offset: 0,
3203 user_data: UserDefinedData::new_zeroed(),
3204 },
3205 connection_id: 0,
3206 event_flag: 0,
3207 flags: Default::default(),
3208 },
3209 );
3210 server.send(in_msg(
3211 MessageType::OPEN_CHANNEL_RESULT,
3212 protocol::OpenResult {
3213 channel_id: ChannelId(0),
3214 open_id: 0,
3215 status: protocol::STATUS_SUCCESS as u32,
3216 },
3217 ));
3218 };
3219
3220 (open, server_open).join().await.0.unwrap();
3221
3222 drop(channel);
3224
3225 check_message(
3226 server.next().await.unwrap(),
3227 protocol::CloseChannel {
3228 channel_id: ChannelId(0),
3229 },
3230 );
3231
3232 server.send(in_msg(
3233 MessageType::RESCIND_CHANNEL_OFFER,
3234 protocol::RescindChannelOffer {
3235 channel_id: ChannelId(0),
3236 },
3237 ));
3238
3239 check_message(
3241 server.next().await.unwrap(),
3242 protocol::RelIdReleased {
3243 channel_id: ChannelId(0),
3244 },
3245 );
3246
3247 let offer = protocol::OfferChannel {
3248 interface_id: Guid::new_random(),
3249 instance_id: Guid::new_random(),
3250 rsvd: [0; 4],
3251 flags: OfferFlags::new(),
3252 mmio_megabytes: 0,
3253 user_defined: UserDefinedData::new_zeroed(),
3254 subchannel_index: 0,
3255 mmio_megabytes_optional: 0,
3256 channel_id: ChannelId(0),
3257 monitor_id: 0,
3258 monitor_allocated: 0,
3259 is_dedicated: 0,
3260 connection_id: 0,
3261 };
3262
3263 server.send(in_msg(MessageType::OFFER_CHANNEL, offer));
3264
3265 connection.offer_recv.next().await.unwrap();
3267 }
3268}