Skip to main content

vmbus_core/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4#![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
26/// The standard, non-redirected synthetic interrupt used by VMBus.
27pub const VMBUS_SINT: u8 = 2;
28
29/// Represents information about a negotiated version.
30#[derive(Copy, Clone, Debug, PartialEq, Eq, Inspect)]
31pub struct VersionInfo {
32    pub version: protocol::Version,
33    pub feature_flags: protocol::FeatureFlags,
34}
35
36/// Represents a constraint on the version or features allowed.
37#[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
61/// Parses a string of the form "major.minor" (e.g "5.3") into a vmbus version number.
62///
63/// N.B. This doesn't check whether the specified version actually exists.
64pub 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
80/// Represents a vmbus message to be sent using the synic.
81impl OutgoingMessage {
82    /// Creates a new `OutgoingMessage` for the specified protocol message.
83    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    /// Creates a new `OutgoingMessage` for the specified protocol message, including additional
97    /// data at the end of the message.
98    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    /// Converts an existing binary message to an `OutgoingMessage`. The slice
111    /// is assumed to contain a valid message.
112    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    /// Gets the binary representation of the message.
125    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/// A request from the guest to connect to the specified hvsocket endpoint.
141#[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/// A notification from the host that a connection request has been handled.
173#[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    /// Create a new result using the service and endpoint ID from the specified request.
182    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}