microsoft/openvmm

Public

mirrored from https://github.com/microsoft/openvmmAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
release/2505-fork

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

openhcl/diag_server/src/lib.rs

265lines · modecode

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! Underhill diagnostics server.
5
6#![cfg(target_os = "linux")]
7
8mod diag_service;
9mod new_pty;
10
11pub use diag_service::DiagRequest;
12pub use diag_service::StartParams;
13
14use anyhow::Context;
15use futures::AsyncWriteExt;
16use futures::FutureExt;
17use mesh::CancelReason;
18use mesh_rpc::server::RpcReceiver;
19use mesh_rpc::service::Code;
20use mesh_rpc::service::Status;
21use pal_async::driver::Driver;
22use pal_async::interest::PollEvents;
23use pal_async::socket::PollReadyExt;
24use pal_async::socket::PolledSocket;
25use pal_async::task::Spawn;
26use pal_async::task::Task;
27use parking_lot::Mutex;
28use socket2::Socket;
29use std::collections::HashMap;
30use std::path::Path;
31use std::pin::pin;
32use std::sync::Arc;
33use unix_socket::UnixListener;
34use vmsocket::VmAddress;
35use vmsocket::VmListener;
36
37/// The diagnostics server, which is a ttrpc server listening on `AF_VSOCK` at
38/// for control and data.
39pub struct DiagServer {
40 // control listener
41 control_listener: Socket,
42 // data listener
43 data_listener: Socket,
44 inner: Arc<Inner>,
45 server: mesh_rpc::Server,
46}
47
48impl DiagServer {
49 /// Creates a server over VmSockets and starts listening.
50 pub fn new_vsock(control_address: VmAddress, data_address: VmAddress) -> anyhow::Result<Self> {
51 tracing::info!(?control_address, "control starting");
52 let control_listener =
53 VmListener::bind(control_address).context("failed to bind socket")?;
54
55 tracing::info!(?data_address, "data starting");
56 let data_listener = VmListener::bind(data_address).context("failed to bind socket")?;
57
58 Ok(Self::new_generic(
59 control_listener.into(),
60 data_listener.into(),
61 ))
62 }
63
64 /// Creates a server over Unix sockets and starts listening.
65 pub fn new_unix(control_address: &Path, data_address: &Path) -> anyhow::Result<Self> {
66 tracing::info!(?control_address, "control starting");
67 let control_listener =
68 UnixListener::bind(control_address).context("failed to bind socket")?;
69
70 tracing::info!(?data_address, "data starting");
71 let data_listener = UnixListener::bind(data_address).context("failed to bind socket")?;
72
73 Ok(Self::new_generic(
74 control_listener.into(),
75 data_listener.into(),
76 ))
77 }
78
79 fn new_generic(control_listener: Socket, data_listener: Socket) -> Self {
80 Self {
81 control_listener,
82 data_listener,
83 server: mesh_rpc::Server::new(),
84 inner: Arc::new(Inner {
85 connections: Mutex::new(DataConnections {
86 next_id: 1, // connection IDs start at 1, as 0 is an invalid ID.
87 active: Default::default(),
88 }),
89 }),
90 }
91 }
92
93 /// Serves requests until `cancel` is dropped.
94 pub async fn serve(
95 mut self,
96 driver: &(impl Driver + Spawn + Clone),
97 cancel: mesh::OneshotReceiver<()>,
98 request_send: mesh::Sender<DiagRequest>,
99 ) -> anyhow::Result<()> {
100 // Disable all diag requests for CVMs. Inspect filtering will be handled
101 // internally more granularly.
102 let (diag_recv, diag2_recv) = if underhill_confidentiality::confidential_filtering_enabled()
103 {
104 (RpcReceiver::disconnected(), RpcReceiver::disconnected())
105 } else {
106 (
107 self.server.add_service::<diag_proto::UnderhillDiag>(),
108 self.server.add_service::<diag_proto::OpenhclDiag>(),
109 )
110 };
111
112 let inspect_recv = self.server.add_service::<inspect_proto::InspectService>();
113
114 // TODO: split the profiler to a separate service provider.
115 let profile_recv = self
116 .server
117 .add_service::<azure_profiler_proto::AzureProfiler>();
118
119 let diag_service = Arc::new(diag_service::DiagServiceHandler::new(
120 request_send,
121 self.inner.clone(),
122 ));
123 let process = diag_service.process_requests(
124 driver,
125 diag_recv,
126 diag2_recv,
127 inspect_recv,
128 profile_recv,
129 );
130
131 let serve = self.server.run(driver, self.control_listener, cancel);
132 let data_connections = self
133 .inner
134 .process_data_connections(driver, self.data_listener);
135
136 futures::future::try_join3(serve, process, data_connections).await?;
137 Ok(())
138 }
139}
140
141#[derive(Debug)]
142struct DataConnectionEntry {
143 /// Sender used to notify the hangup task to return the socket.
144 sender: mesh::OneshotSender<()>,
145 /// Task used to wait for hangup notifications or a request to return the socket.
146 task: Task<Option<PolledSocket<Socket>>>,
147}
148
149#[derive(Debug, Default)]
150struct DataConnections {
151 next_id: u64,
152 active: HashMap<u64, DataConnectionEntry>,
153}
154
155impl DataConnections {
156 fn take_connection(&mut self, id: u64) -> anyhow::Result<DataConnectionEntry> {
157 self.active
158 .remove(&id)
159 .ok_or_else(|| anyhow::anyhow!("invalid connection id"))
160 }
161}
162
163struct Inner {
164 connections: Mutex<DataConnections>,
165}
166
167impl Inner {
168 async fn take_connection(&self, id: u64) -> anyhow::Result<PolledSocket<Socket>> {
169 let DataConnectionEntry { sender, task } = self.connections.lock().take_connection(id)?;
170
171 sender.send(());
172 task.await
173 .ok_or_else(|| anyhow::anyhow!("connection disconnected"))
174 }
175
176 /// Listen for data connections and add them to the internal connections lookup table as they arrive.
177 async fn process_data_connections(
178 self: &Arc<Self>,
179 driver: &(impl Driver + Spawn + Clone),
180 listener: Socket,
181 ) -> anyhow::Result<()> {
182 let mut listener = PolledSocket::new(driver, listener)?;
183
184 loop {
185 let (connection, _addr) = listener.accept().await?;
186 let mut socket = PolledSocket::new(driver, connection)?;
187 let inner = Arc::downgrade(self);
188
189 // Send the 8 byte connection id, then stash the connection in the lookup table to be used later.
190 let id;
191 {
192 let mut state = self.connections.lock();
193 id = state.next_id;
194 state.next_id += 1;
195
196 tracing::debug!(id, "new data connection");
197 }
198
199 let (sender, recv) = mesh::oneshot();
200
201 // Spawn a task that returns the socket when asked to, or removes itself from the map if disconnected.
202 let task = driver.spawn(format!("data connection {} waiting", id), async move {
203 match socket.write_all(&id.to_ne_bytes()).await {
204 Ok(_) => {}
205 Err(error) => {
206 tracing::trace!(?error, "error writing connection id, removing.");
207 if let Some(state) = inner.upgrade() {
208 state.connections.lock().active.remove(&id);
209 }
210
211 return None;
212 }
213 }
214
215 let mut return_future = pin!(async { recv.await.is_ok() }.fuse());
216 let hangup = futures::select! { // race semantics
217 _ = socket.wait_ready(PollEvents::RDHUP).fuse() => true,
218 _ = return_future => false,
219 };
220
221 if hangup {
222 // Other side has disconnected, remove from the table if not already done.
223 tracing::trace!(id, "data connection disconnected");
224 if let Some(state) = inner.upgrade() {
225 state.connections.lock().active.remove(&id);
226 }
227
228 None
229 } else {
230 Some(socket)
231 }
232 });
233
234 let mut state = self.connections.lock();
235 let result = state
236 .active
237 .insert(id, DataConnectionEntry { sender, task });
238
239 if result.is_some() {
240 anyhow::bail!("connection id reused");
241 }
242 }
243 }
244}
245
246fn grpc_result<T>(result: Result<anyhow::Result<T>, CancelReason>) -> Result<T, Status> {
247 match result {
248 Ok(result) => match result {
249 Ok(value) => Ok(value),
250 Err(err) => Err(Status {
251 code: Code::Unknown as i32,
252 message: format!("{:#}", err),
253 details: vec![],
254 }),
255 },
256 Err(err) => Err(Status {
257 code: match &err {
258 CancelReason::Cancelled => Code::Cancelled,
259 CancelReason::DeadlineExceeded => Code::DeadlineExceeded,
260 } as i32,
261 message: format!("{:#}", err),
262 details: vec![],
263 }),
264 }
265}
266