1#![cfg_attr(not(test), no_std)]
5#![expect(missing_docs)]
6#![forbid(unsafe_code)]
7
8extern crate alloc;
9
10pub mod protocol;
11
12use alloc::format;
13use alloc::string::String;
14use core::str::FromStr;
15use guid::Guid;
16use inspect::Inspect;
17use protocol::HEADER_SIZE;
18use protocol::MAX_MESSAGE_SIZE;
19use protocol::MessageHeader;
20use protocol::VmbusMessage;
21use thiserror::Error;
22use zerocopy::Immutable;
23use zerocopy::IntoBytes;
24use zerocopy::KnownLayout;
25
26pub const VMBUS_SINT: u8 = 2;
28
29#[derive(Copy, Clone, Debug, PartialEq, Eq, Inspect)]
31pub struct VersionInfo {
32 pub version: protocol::Version,
33 pub feature_flags: protocol::FeatureFlags,
34}
35
36#[derive(Copy, Clone, Debug)]
38pub struct MaxVersionInfo {
39 pub version: u32,
40 pub feature_flags: protocol::FeatureFlags,
41}
42
43impl MaxVersionInfo {
44 pub fn new(version: u32) -> Self {
45 Self {
46 version,
47 feature_flags: protocol::FeatureFlags::new(),
48 }
49 }
50}
51
52impl From<VersionInfo> for MaxVersionInfo {
53 fn from(info: VersionInfo) -> Self {
54 Self {
55 version: info.version as u32,
56 feature_flags: info.feature_flags,
57 }
58 }
59}
60
61pub fn parse_vmbus_version(value: &str) -> Result<u32, String> {
65 || -> Option<u32> {
66 let (major, minor) = value.split_once('.')?;
67 let major = u16::from_str(major).ok()?;
68 let minor = u16::from_str(minor).ok()?;
69 Some(protocol::make_version(major, minor))
70 }()
71 .ok_or_else(|| format!("invalid vmbus version '{}'", value))
72}
73
74#[derive(Clone, Debug)]
75pub struct OutgoingMessage {
76 data: [u8; MAX_MESSAGE_SIZE],
77 len: u8,
78}
79
80impl OutgoingMessage {
82 pub fn new<T: IntoBytes + Immutable + KnownLayout + VmbusMessage>(message: &T) -> Self {
84 let mut data = [0; MAX_MESSAGE_SIZE];
85 let header = MessageHeader::new(T::MESSAGE_TYPE);
86 let message_bytes = message.as_bytes();
87 let len = HEADER_SIZE + message_bytes.len();
88 data[..HEADER_SIZE].copy_from_slice(header.as_bytes());
89 data[HEADER_SIZE..len].copy_from_slice(message_bytes);
90 Self {
91 data,
92 len: len as u8,
93 }
94 }
95
96 pub fn with_data<T: IntoBytes + Immutable + KnownLayout + VmbusMessage>(
99 message: &T,
100 data: &[u8],
101 ) -> Self {
102 let mut message = OutgoingMessage::new(message);
103 let old_len = message.len as usize;
104 let len = old_len + data.len();
105 message.data[old_len..len].copy_from_slice(data);
106 message.len = len as u8;
107 message
108 }
109
110 pub fn from_message(message: &[u8]) -> Result<Self, MessageTooLarge> {
113 if message.len() > MAX_MESSAGE_SIZE {
114 return Err(MessageTooLarge);
115 }
116 let mut data = [0; MAX_MESSAGE_SIZE];
117 data[0..message.len()].copy_from_slice(message);
118 Ok(Self {
119 data,
120 len: message.len() as u8,
121 })
122 }
123
124 pub fn data(&self) -> &[u8] {
126 &self.data[..self.len as usize]
127 }
128}
129
130impl PartialEq for OutgoingMessage {
131 fn eq(&self, other: &Self) -> bool {
132 self.len == other.len && self.data[..self.len as usize] == other.data[..self.len as usize]
133 }
134}
135
136#[derive(Debug, Error)]
137#[error("a synic message exceeds the maximum length")]
138pub struct MessageTooLarge;
139
140#[derive(Copy, Clone, Debug, Hash, Eq, PartialEq, Inspect)]
142pub struct HvsockConnectRequest {
143 pub service_id: Guid,
144 pub endpoint_id: Guid,
145 pub silo_id: Guid,
146 pub hosted_silo_unaware: bool,
147}
148
149impl HvsockConnectRequest {
150 pub fn from_message(value: protocol::TlConnectRequest2, hosted_silo_unaware: bool) -> Self {
151 Self {
152 service_id: value.base.service_id,
153 endpoint_id: value.base.endpoint_id,
154 silo_id: value.silo_id,
155 hosted_silo_unaware,
156 }
157 }
158}
159
160impl From<HvsockConnectRequest> for protocol::TlConnectRequest2 {
161 fn from(value: HvsockConnectRequest) -> Self {
162 Self {
163 base: protocol::TlConnectRequest {
164 endpoint_id: value.endpoint_id,
165 service_id: value.service_id,
166 },
167 silo_id: value.silo_id,
168 }
169 }
170}
171
172#[derive(Copy, Clone, Debug, Hash, Eq, PartialEq)]
174pub struct HvsockConnectResult {
175 pub service_id: Guid,
176 pub endpoint_id: Guid,
177 pub success: bool,
178}
179
180impl HvsockConnectResult {
181 pub fn from_request(request: &HvsockConnectRequest, success: bool) -> Self {
183 Self {
184 service_id: request.service_id,
185 endpoint_id: request.endpoint_id,
186 success,
187 }
188 }
189}
190
191impl From<protocol::TlConnectResult> for HvsockConnectResult {
192 fn from(value: protocol::TlConnectResult) -> Self {
193 Self {
194 service_id: value.service_id,
195 endpoint_id: value.endpoint_id,
196 success: value.status == protocol::STATUS_SUCCESS,
197 }
198 }
199}
200
201#[cfg(test)]
202mod tests {
203 use super::*;
204 use crate::protocol::ChannelId;
205 use crate::protocol::GpadlId;
206
207 #[test]
208 fn test_outgoing_message() {
209 let message = OutgoingMessage::new(&protocol::CloseChannel {
210 channel_id: ChannelId(5),
211 });
212
213 assert_eq!(&[0x7, 0, 0, 0, 0, 0, 0, 0, 0x5, 0, 0, 0], message.data())
214 }
215
216 #[test]
217 fn test_outgoing_message_empty() {
218 let message = OutgoingMessage::new(&protocol::Unload {});
219
220 assert_eq!(&[0x10, 0, 0, 0, 0, 0, 0, 0], message.data())
221 }
222
223 #[test]
224 fn test_outgoing_message_with_data() {
225 let message = OutgoingMessage::with_data(
226 &protocol::GpadlHeader {
227 channel_id: ChannelId(5),
228 gpadl_id: GpadlId(1),
229 len: 7,
230 count: 6,
231 },
232 &[0xa, 0xb, 0xc, 0xd],
233 );
234
235 assert_eq!(
236 &[
237 0x8, 0, 0, 0, 0, 0, 0, 0, 0x5, 0, 0, 0, 0x1, 0, 0, 0, 0x7, 0, 0x6, 0, 0xa, 0xb,
238 0xc, 0xd
239 ],
240 message.data()
241 )
242 }
243}