1#![cfg(target_os = "linux")]
11#![forbid(unsafe_code)]
12
13use disk_backend::DiskError;
14use disk_backend::DiskIo;
15use guest_emulation_transport::GuestEmulationTransportClient;
16use guestmem::MemoryRead;
17use guestmem::MemoryWrite;
18use inspect::Inspect;
19use save_restore::SavedBlockStorageMetadata;
20use scsi_buffers::RequestBuffers;
21use std::io;
22use thiserror::Error;
23
24#[derive(Clone, Debug, Inspect)]
26pub struct GetVmgsDisk {
27 get: GuestEmulationTransportClient,
28 sector_size: u32,
29 sector_shift: u32,
30 physical_sector_size: u32,
31 sector_count: u64,
32 max_transfer_sectors: u32,
33 max_transfer_size_bytes: u32,
34}
35
36#[derive(Debug, Error)]
38pub enum NewGetVmgsDiskError {
39 #[error("GET VMGS IO error")]
41 Io(#[source] guest_emulation_transport::error::VmgsIoError),
42 #[error("invalid sector size")]
44 InvalidSectorSize,
45 #[error("invalid physical sector size")]
47 InvalidPhysicalSectorSize,
48 #[error("invalid sector count")]
50 InvalidSectorCount,
51 #[error("disk ends with a partial physical sector")]
53 IncompletePhysicalSector,
54 #[error("transfer size is smaller than the physical sector size")]
56 InvalidMaxTransferSize,
57}
58
59impl GetVmgsDisk {
60 pub async fn new(get: GuestEmulationTransportClient) -> Result<Self, NewGetVmgsDiskError> {
63 let response = get
64 .vmgs_get_device_info()
65 .await
66 .map_err(NewGetVmgsDiskError::Io)?;
67 Self::new_inner(
68 get,
69 response.bytes_per_logical_sector.into(),
70 response.bytes_per_physical_sector.into(),
71 response.capacity,
72 response.maximum_transfer_size_bytes,
73 )
74 }
75
76 pub fn restore_with_meta(
88 get: GuestEmulationTransportClient,
89 meta: SavedBlockStorageMetadata,
90 ) -> Result<Self, NewGetVmgsDiskError> {
91 Self::new_inner(
92 get,
93 meta.sector_size,
94 meta.physical_sector_size,
95 meta.sector_count,
96 meta.max_transfer_size_bytes,
97 )
98 }
99
100 pub fn save_meta(&self) -> SavedBlockStorageMetadata {
103 SavedBlockStorageMetadata {
104 capacity: self.sector_count * self.sector_size as u64,
105 logical_sector_size: self.sector_size,
106 sector_count: self.sector_count,
107 sector_size: self.sector_size,
108 physical_sector_size: self.physical_sector_size,
109 max_transfer_size_bytes: self.max_transfer_size_bytes,
110 }
111 }
112
113 fn new_inner(
114 get: GuestEmulationTransportClient,
115 sector_size: u32,
116 physical_sector_size: u32,
117 sector_count: u64,
118 max_transfer_size: u32,
119 ) -> Result<Self, NewGetVmgsDiskError> {
120 if !sector_size.is_power_of_two() {
121 Err(NewGetVmgsDiskError::InvalidSectorSize)
122 } else if !physical_sector_size.is_power_of_two() || physical_sector_size < sector_size {
123 Err(NewGetVmgsDiskError::InvalidPhysicalSectorSize)
124 } else if sector_count.checked_mul(sector_size as u64).is_none() {
125 Err(NewGetVmgsDiskError::InvalidSectorCount)
126 } else if !sector_count.is_multiple_of((physical_sector_size / sector_size) as u64) {
127 Err(NewGetVmgsDiskError::IncompletePhysicalSector)
128 } else if max_transfer_size < physical_sector_size {
129 Err(NewGetVmgsDiskError::InvalidMaxTransferSize)
130 } else {
131 Ok(GetVmgsDisk {
132 get,
133 sector_size,
134 sector_shift: sector_size.trailing_zeros(),
135 physical_sector_size,
136 sector_count,
137 max_transfer_size_bytes: max_transfer_size,
138 max_transfer_sectors: max_transfer_size / physical_sector_size
140 * physical_sector_size
141 / sector_size,
142 })
143 }
144 }
145}
146
147impl DiskIo for GetVmgsDisk {
148 fn disk_type(&self) -> &str {
149 "vmgs-get"
150 }
151
152 fn sector_count(&self) -> u64 {
153 self.sector_count
154 }
155
156 fn sector_size(&self) -> u32 {
157 self.sector_size
158 }
159
160 fn disk_id(&self) -> Option<[u8; 16]> {
161 None
162 }
163
164 fn physical_sector_size(&self) -> u32 {
165 self.physical_sector_size
166 }
167
168 fn is_fua_respected(&self) -> bool {
169 false
170 }
171
172 fn is_read_only(&self) -> bool {
173 false
174 }
175
176 async fn read_vectored(
177 &self,
178 buffers: &RequestBuffers<'_>,
179 mut sector: u64,
180 ) -> Result<(), DiskError> {
181 let mut writer = buffers.writer();
182 let mut remaining_sectors = buffers.len() >> self.sector_shift;
183 if sector + remaining_sectors as u64 > self.sector_count {
184 return Err(DiskError::IllegalBlock);
185 }
186 while remaining_sectors != 0 {
187 let this_sector_count = remaining_sectors.min(self.max_transfer_sectors as usize);
188 let data = self
189 .get
190 .vmgs_read(sector, this_sector_count as u32, self.sector_size)
191 .await
192 .map_err(|err| DiskError::Io(io::Error::other(err)))?;
193
194 writer.write(&data)?;
195 sector += this_sector_count as u64;
196 remaining_sectors -= this_sector_count;
197 }
198 Ok(())
199 }
200
201 async fn write_vectored(
202 &self,
203 buffers: &RequestBuffers<'_>,
204 mut sector: u64,
205 _fua: bool,
206 ) -> Result<(), DiskError> {
207 let mut reader = buffers.reader();
208 let mut remaining_sector_count = buffers.len() >> self.sector_shift;
209 if sector + remaining_sector_count as u64 > self.sector_count {
210 return Err(DiskError::IllegalBlock);
211 }
212 while remaining_sector_count != 0 {
213 let this_sector_count = remaining_sector_count.min(self.max_transfer_sectors as usize);
214 let data = reader.read_n(this_sector_count << self.sector_shift)?;
215 self.get
216 .vmgs_write(sector, data, self.sector_size)
217 .await
218 .map_err(|err| DiskError::Io(io::Error::other(err)))?;
219
220 remaining_sector_count -= this_sector_count;
221 sector += this_sector_count as u64;
222 }
223 Ok(())
224 }
225
226 async fn sync_cache(&self) -> Result<(), DiskError> {
228 self.get
229 .vmgs_flush()
230 .await
231 .map_err(|err| DiskError::Io(io::Error::other(err)))
232 }
233
234 async fn unmap(
235 &self,
236 _sector: u64,
237 _count: u64,
238 _block_level_only: bool,
239 ) -> Result<(), DiskError> {
240 Ok(())
241 }
242
243 fn unmap_behavior(&self) -> disk_backend::UnmapBehavior {
244 disk_backend::UnmapBehavior::Ignored
245 }
246}
247
248pub mod save_restore {
250 use mesh::payload::Protobuf;
251
252 #[derive(Protobuf, Clone)]
254 #[mesh(package = "vmgs")]
255 pub struct SavedBlockStorageMetadata {
256 #[mesh(1)]
258 pub capacity: u64,
259 #[mesh(2)]
261 pub logical_sector_size: u32,
262 #[mesh(3)]
264 pub sector_count: u64,
265 #[mesh(4)]
267 pub sector_size: u32,
268 #[mesh(5)]
270 pub physical_sector_size: u32,
271 #[mesh(6)]
273 pub max_transfer_size_bytes: u32,
274 }
275}
276
277#[cfg(test)]
279mod tests {
280 use super::*;
281 use disk_backend::Disk;
282 use guest_emulation_transport::api::ProtocolVersion;
283 use guest_emulation_transport::test_utilities::TestGet;
284 use guest_emulation_transport::test_utilities::new_transport_pair;
285 use pal_async::DefaultDriver;
286 use pal_async::async_test;
287 use pal_async::task::Task;
288 use vmgs::FileId;
289 use vmgs::Vmgs;
290 use vmgs_broker::VmgsClient;
291 use vmgs_broker::spawn_vmgs_broker;
292
293 async fn spawn_vmgs(driver: &DefaultDriver) -> (VmgsClient, TestGet, Task<()>) {
294 let get = new_transport_pair(driver, None, ProtocolVersion::NICKEL_REV2, None, None).await;
295 let vmgs_get = GetVmgsDisk::new(get.client.clone()).await.unwrap();
296 let vmgs = Vmgs::format_new(Disk::new(vmgs_get).unwrap(), None)
297 .await
298 .unwrap();
299 let (vmgs, task) = spawn_vmgs_broker(driver, vmgs);
300 (vmgs, get, task)
301 }
302
303 #[async_test]
305 async fn sector_range_conformance(driver: DefaultDriver) {
306 let get = new_transport_pair(&driver, None, ProtocolVersion::NICKEL_REV2, None, None).await;
307 let disk = Disk::new(GetVmgsDisk::new(get.client.clone()).await.unwrap()).unwrap();
308 storage_tests::sector_range::test_disk_sector_range_conformance(&disk).await;
309 }
310
311 #[async_test]
312 async fn basic_read_write(driver: DefaultDriver) {
313 let (vmgs, _get, _task) = spawn_vmgs(&driver).await;
314 let file_id = FileId::BIOS_NVRAM;
315
316 let buf = b"hello world".to_vec();
318 vmgs.write_file(file_id, buf.clone()).await.unwrap();
319
320 let info = vmgs.get_file_info(file_id).await.unwrap();
322 assert_eq!(info.valid_bytes as usize, buf.len());
323 let read_buf = vmgs.read_file(file_id).await.unwrap();
324
325 assert_eq!(buf, read_buf);
326 }
327
328 #[async_test]
329 async fn multiple_read_write(driver: DefaultDriver) {
330 let (vmgs, _get, _task) = spawn_vmgs(&driver).await;
331 let file_id_1 = FileId::BIOS_NVRAM;
332 let file_id_2 = FileId::TPM_PPI;
333 let buf_1 = b"Data data data".to_vec();
334 let buf_2 = b"password".to_vec();
335 let buf_3 = b"other data data".to_vec();
336
337 vmgs.write_file(file_id_1, buf_1.clone()).await.unwrap();
338 let info = vmgs.get_file_info(file_id_1).await.unwrap();
339 assert_eq!(info.valid_bytes as usize, buf_1.len());
340 let read_buf_1 = vmgs.read_file(file_id_1).await.unwrap();
341 assert_eq!(buf_1, read_buf_1);
342
343 vmgs.write_file(file_id_2, buf_2.clone()).await.unwrap();
344 let info = vmgs.get_file_info(file_id_2).await.unwrap();
345 assert_eq!(info.valid_bytes as usize, buf_2.len());
346 let read_buf_2 = vmgs.read_file(file_id_2).await.unwrap();
347 assert_eq!(buf_2, read_buf_2);
348
349 vmgs.write_file(file_id_1, buf_3.clone()).await.unwrap();
350 let info = vmgs.get_file_info(file_id_1).await.unwrap();
351 assert_eq!(info.valid_bytes as usize, buf_3.len());
352 let read_buf_3 = vmgs.read_file(file_id_1).await.unwrap();
353 assert_eq!(buf_3, read_buf_3);
354
355 vmgs.write_file(file_id_1, buf_1.clone()).await.unwrap();
356 let info = vmgs.get_file_info(file_id_1).await.unwrap();
357 assert_eq!(info.valid_bytes as usize, buf_1.len());
358 let read_buf_1 = vmgs.read_file(file_id_1).await.unwrap();
359 assert_eq!(buf_1, read_buf_1);
360
361 vmgs.write_file(file_id_2, buf_2.clone()).await.unwrap();
362 let info = vmgs.get_file_info(file_id_2).await.unwrap();
363 assert_eq!(info.valid_bytes as usize, buf_2.len());
364 let read_buf_2 = vmgs.read_file(file_id_2).await.unwrap();
365 assert_eq!(buf_2, read_buf_2);
366
367 vmgs.write_file(file_id_1, buf_3.clone()).await.unwrap();
368 let info = vmgs.get_file_info(file_id_1).await.unwrap();
369 assert_eq!(info.valid_bytes as usize, buf_3.len());
370 let read_buf_3 = vmgs.read_file(file_id_1).await.unwrap();
371 assert_eq!(buf_3, read_buf_3);
372 }
373
374 #[async_test]
375 async fn test_empty_write(driver: DefaultDriver) {
376 let (vmgs, _get, _task) = spawn_vmgs(&driver).await;
377 let file_id = FileId::BIOS_NVRAM;
378
379 let buf: Vec<u8> = Vec::new();
380 vmgs.write_file(file_id, buf.clone()).await.unwrap();
381
382 let info = vmgs.get_file_info(file_id).await.unwrap();
384 assert_eq!(info.valid_bytes as usize, 0);
385 let read_buf = vmgs.read_file(file_id).await.unwrap();
386
387 assert_eq!(buf, read_buf);
388 assert_eq!(read_buf.len(), 0);
389 }
390
391 #[async_test]
392 async fn test_read_write_large(driver: DefaultDriver) {
393 let (vmgs, _get, _task) = spawn_vmgs(&driver).await;
394 let file_id = FileId::BIOS_NVRAM;
395
396 let buf: Vec<u8> = (0..).map(|x| x as u8).take(1024 * 4 * 4 + 1).collect();
398 vmgs.write_file(file_id, buf.clone()).await.unwrap();
399
400 let info = vmgs.get_file_info(file_id).await.unwrap();
402 assert_eq!(info.valid_bytes as usize, buf.len());
403 let read_buf = vmgs.read_file(file_id).await.unwrap();
404
405 assert_eq!(buf, read_buf);
406 }
407
408 #[async_test]
409 async fn test_read_write_encryption(driver: DefaultDriver) {
410 let get = new_transport_pair(&driver, None, ProtocolVersion::NICKEL_REV2, None, None).await;
411 let vmgs_get = GetVmgsDisk::new(get.client.clone()).await.unwrap();
412 let mut vmgs = Vmgs::format_new(Disk::new(vmgs_get).unwrap(), None)
413 .await
414 .unwrap();
415 let file_id = FileId::BIOS_NVRAM;
416 let encryption_key = [1; 32];
417
418 vmgs.update_encryption_key(&encryption_key, vmgs::EncryptionAlgorithm::AES_GCM)
419 .await
420 .unwrap();
421
422 let buf: Vec<u8> = (0..).map(|x| x as u8).take(1024 * 4 * 4 + 1).collect();
424 vmgs.write_file_encrypted(file_id, &buf).await.unwrap();
425
426 let info = vmgs.get_file_info(file_id).unwrap();
428 assert_eq!(info.valid_bytes as usize, buf.len());
429 let read_buf = vmgs.read_file(file_id).await.unwrap();
430
431 assert_eq!(buf, read_buf);
432
433 drop(vmgs);
434
435 let vmgs_get = GetVmgsDisk::new(get.client.clone()).await.unwrap();
436 let mut vmgs = Vmgs::open(Disk::new(vmgs_get).unwrap(), None)
437 .await
438 .unwrap();
439
440 let read_buf = vmgs.read_file_raw(file_id).await.unwrap();
441
442 assert_ne!(buf, read_buf);
443
444 vmgs.unlock_with_encryption_key(&encryption_key)
445 .await
446 .unwrap();
447
448 let (vmgs, _task) = spawn_vmgs_broker(&driver, vmgs);
449
450 let info = vmgs.get_file_info(file_id).await.unwrap();
451 assert_eq!(info.valid_bytes as usize, buf.len());
452 let read_buf = vmgs.read_file(file_id).await.unwrap();
453
454 assert_eq!(buf, read_buf);
455 }
456}