1use std::array;
8use std::ops::Deref;
9use std::ops::DerefMut;
10use std::sync::atomic::AtomicU64;
11use std::sync::atomic::Ordering;
12use x86defs::snp::SecureAvicControl;
13use x86defs::snp::SevEventInjectInfo;
14use x86defs::snp::SevFeatures;
15use x86defs::snp::SevSelector;
16use x86defs::snp::SevVirtualInterruptControl;
17use x86defs::snp::SevVmsa;
18use x86defs::snp::SevXmmRegister;
19use zerocopy::FromZeros;
20use zerocopy::IntoBytes;
21
22pub struct VmsaWrapper<'a, T> {
24 vmsa: T,
25 bitmap: &'a [u8; 64],
26}
27
28impl<'a, T> VmsaWrapper<'a, T> {
29 pub(crate) fn new(vmsa: T, bitmap: &'a [u8; 64]) -> Self {
31 VmsaWrapper { vmsa, bitmap }
32 }
33}
34
35impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
37 fn get_u64(&self, offset: usize) -> u64 {
39 assert!(offset.is_multiple_of(8));
40 let vmsa_raw = &self.vmsa;
41 let v = u64::from_ne_bytes(vmsa_raw.as_bytes()[offset..offset + 8].try_into().unwrap());
42 if is_protected(self.bitmap, offset) {
43 v ^ self.vmsa.register_protection_nonce
44 } else {
45 v
46 }
47 }
48 fn get_u32(&self, offset: usize) -> u32 {
50 assert!(offset.is_multiple_of(4));
51 (self.get_u64(offset & !7) >> ((offset & 4) * 8)) as u32
52 }
53 fn get_u128(&self, offset: usize) -> u128 {
55 self.get_u64(offset) as u128 | ((self.get_u64(offset + 8) as u128) << 64)
56 }
57
58 pub fn xmm_registers(&self, n: usize) -> u128 {
60 assert!(n < 16);
61 let off = std::mem::offset_of!(SevVmsa, xmm_registers) + (n * 16);
62 self.get_u128(off)
63 }
64
65 pub fn ymm_registers(&self, n: usize) -> u128 {
67 assert!(n < 16);
68 let off = std::mem::offset_of!(SevVmsa, ymm_registers) + (n * 16);
69 self.get_u128(off)
70 }
71
72 pub fn x87_registers(&self) -> [u64; 10] {
74 let base = std::mem::offset_of!(SevVmsa, x87_registers);
75 array::from_fn(|i| i * 8).map(|offset| self.get_u64(base + offset))
76 }
77}
78
79impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
81 fn set_u64(&self, v: u64, offset: usize) -> u64 {
83 assert!(offset.is_multiple_of(8));
84 if is_protected(self.bitmap, offset) {
85 v ^ self.vmsa.register_protection_nonce
86 } else {
87 v
88 }
89 }
90 fn set_u32(&self, v: u32, offset: usize) -> u32 {
92 assert!(offset.is_multiple_of(4));
93 let val = (v as u64) << ((offset & 4) * 8);
94 (self.set_u64(val, offset & !7) >> ((offset & 4) * 8)) as u32
95 }
96 fn set_u128(&self, v: u128, offset: usize) -> u128 {
98 self.set_u64(v as u64, offset) as u128
99 | ((self.set_u64((v >> 64) as u64, offset + 8) as u128) << 64)
100 }
101
102 pub fn reset(&mut self, vmsa_reg_prot: bool) {
104 *self.vmsa = FromZeros::new_zeroed();
105 if vmsa_reg_prot {
106 getrandom::fill(self.vmsa.register_protection_nonce.as_mut_bytes())
108 .expect("rng failure");
109 let nonce = self.vmsa.register_protection_nonce;
110 let chunk_size = 8;
111 for (i, b) in self
112 .vmsa
113 .as_mut_bytes()
114 .chunks_exact_mut(chunk_size)
115 .enumerate()
116 {
117 let field_offset = i * chunk_size;
118 if field_offset == (std::mem::offset_of!(SevVmsa, vmpl) & !7)
120 || field_offset == std::mem::offset_of!(SevVmsa, exit_info1)
121 || field_offset == std::mem::offset_of!(SevVmsa, exit_info2)
122 || field_offset == std::mem::offset_of!(SevVmsa, exit_int_info)
123 || field_offset == std::mem::offset_of!(SevVmsa, sev_features)
124 || field_offset == std::mem::offset_of!(SevVmsa, v_intr_cntrl)
125 || field_offset == std::mem::offset_of!(SevVmsa, guest_error_code)
126 || field_offset == std::mem::offset_of!(SevVmsa, virtual_tom)
127 {
128 assert!(!is_protected(self.bitmap, field_offset));
129 }
130 if is_protected(self.bitmap, field_offset) {
131 b.copy_from_slice(&nonce.to_ne_bytes());
132 }
133 }
134 }
135 }
136
137 pub fn set_xmm_registers(&mut self, n: usize, v: u128) {
139 assert!(n < 16);
140 let off = std::mem::offset_of!(SevVmsa, xmm_registers) + (n * 16);
141 let val: SevXmmRegister = self.set_u128(v, off).into();
142 let vmsa_raw = &mut *self.vmsa;
143 vmsa_raw.xmm_registers[n] = val;
144 }
145
146 pub fn set_ymm_registers(&mut self, n: usize, v: u128) {
148 assert!(n < 16);
149 let off = std::mem::offset_of!(SevVmsa, ymm_registers) + (n * 16);
150 let val: SevXmmRegister = self.set_u128(v, off).into();
151 let vmsa_raw = &mut *self.vmsa;
152 vmsa_raw.ymm_registers[n] = val;
153 }
154
155 pub fn set_x87_registers(&mut self, v: &[u64; 10]) {
157 let base = std::mem::offset_of!(SevVmsa, x87_registers);
158 for (i, new_v) in v.iter().enumerate() {
159 let val = self.set_u64(*new_v, base + (i * 8));
160 self.vmsa.x87_registers[i] = val;
161 }
162 }
163
164 pub fn guest_busy_bit_test_and_set(&mut self) -> bool {
166 const VINTR_GUEST_BUSYBIT_MASK: u64 = 1u64 << 63;
167 let prev = unsafe { &*(core::ptr::from_ref(&self.vmsa.v_intr_cntrl).cast::<AtomicU64>()) }
173 .fetch_or(VINTR_GUEST_BUSYBIT_MASK, Ordering::SeqCst);
174 (prev & VINTR_GUEST_BUSYBIT_MASK) != 0
175 }
176}
177
178fn is_protected(bitmap: &[u8; 64], field_offset: usize) -> bool {
180 let byte_index = field_offset / 64;
181 let bit_index = (field_offset % 64) / 8;
182 bitmap[byte_index] & (1 << bit_index) != 0
183}
184
185macro_rules! regss {
186 ($reg:ident, $set:ident) => {
187 impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
188 pub fn $reg(&self) -> SevSelector {
190 SevSelector::from(self.get_u128(std::mem::offset_of!(SevVmsa, $reg)))
191 }
192 }
193 impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
194 pub fn $set(&mut self, v: SevSelector) {
196 let val = SevSelector::from(
197 self.set_u128(v.as_u128(), std::mem::offset_of!(SevVmsa, $reg)),
198 );
199 let vmsa_raw = &mut *self.vmsa;
200 vmsa_raw.$reg = val;
201 }
202 }
203 };
204}
205macro_rules! reg64 {
206 ($reg:ident, $set:ident) => {
207 impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
208 pub fn $reg(&self) -> u64 {
210 self.get_u64(std::mem::offset_of!(SevVmsa, $reg))
211 }
212 }
213 impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
214 pub fn $set(&mut self, v: u64) {
216 let val = self.set_u64(v, std::mem::offset_of!(SevVmsa, $reg));
217 let vmsa_raw = &mut *self.vmsa;
218 vmsa_raw.$reg = val;
219 }
220 }
221 };
222}
223macro_rules! reg32 {
224 ($reg:ident, $set:ident) => {
225 impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
226 pub fn $reg(&self) -> u32 {
228 self.get_u32(std::mem::offset_of!(SevVmsa, $reg))
229 }
230 }
231 impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
232 pub fn $set(&mut self, v: u32) {
234 let val = self.set_u32(v, std::mem::offset_of!(SevVmsa, $reg));
235 let vmsa_raw = &mut *self.vmsa;
236 vmsa_raw.$reg = val;
237 }
238 }
239 };
240}
241macro_rules! get_reg_direct {
242 ($reg:ident, $ty:ty) => {
243 impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
244 pub fn $reg(&self) -> $ty {
246 let vmsa_raw = &self.vmsa;
247 vmsa_raw.$reg
248 }
249 }
250 };
251}
252macro_rules! reg_direct {
253 ($reg:ident, $set:ident, $ty:ty) => {
254 get_reg_direct!($reg, $ty);
255 impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
256 pub fn $set(&mut self, v: $ty) {
258 let vmsa_raw = &mut *self.vmsa;
259 vmsa_raw.$reg = v;
260 }
261 }
262 };
263}
264macro_rules! reg_direct_mut {
265 ($reg:ident, $set:ident, $ty:ty) => {
266 get_reg_direct!($reg, $ty);
267 impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
268 pub fn $set(&mut self) -> &mut $ty {
270 &mut self.vmsa.$reg
271 }
272 }
273 };
274}
275
276reg_direct!(vmpl, set_vmpl, u8);
277get_reg_direct!(cpl, u8);
278get_reg_direct!(exit_info1, u64);
279get_reg_direct!(exit_info2, u64);
280reg_direct!(exit_int_info, set_exit_int_info, u64);
281reg_direct_mut!(sev_features, sev_features_mut, SevFeatures);
282reg_direct_mut!(v_intr_cntrl, v_intr_cntrl_mut, SevVirtualInterruptControl);
283reg_direct!(virtual_tom, set_virtual_tom, u64);
284reg_direct!(event_inject, set_event_inject, SevEventInjectInfo);
285reg_direct!(guest_error_code, set_guest_error_code, u64);
286reg_direct_mut!(
287 secure_avic_control,
288 secure_avic_control_mut,
289 SecureAvicControl
290);
291regss!(es, set_es);
292regss!(cs, set_cs);
293regss!(ss, set_ss);
294regss!(ds, set_ds);
295regss!(fs, set_fs);
296regss!(gs, set_gs);
297regss!(gdtr, set_gdtr);
298regss!(ldtr, set_ldtr);
299regss!(idtr, set_idtr);
300regss!(tr, set_tr);
301reg64!(pl0_ssp, set_pl0_ssp);
302reg64!(pl1_ssp, set_pl1_ssp);
303reg64!(pl2_ssp, set_pl2_ssp);
304reg64!(pl3_ssp, set_pl3_ssp);
305reg64!(u_cet, set_u_cet);
306reg64!(efer, set_efer);
307reg64!(xss, set_xss);
308reg64!(cr4, set_cr4);
309reg64!(cr3, set_cr3);
310reg64!(cr0, set_cr0);
311reg64!(dr7, set_dr7);
312reg64!(dr6, set_dr6);
313reg64!(rflags, set_rflags);
314reg64!(rip, set_rip);
315reg64!(dr0, set_dr0);
316reg64!(dr1, set_dr1);
317reg64!(dr2, set_dr2);
318reg64!(dr3, set_dr3);
319reg64!(rsp, set_rsp);
320reg64!(s_cet, set_s_cet);
321reg64!(ssp, set_ssp);
322reg64!(interrupt_ssp_table_addr, set_interrupt_ssp_table_addr);
323reg64!(rax, set_rax);
324reg64!(star, set_star);
325reg64!(lstar, set_lstar);
326reg64!(cstar, set_cstar);
327reg64!(sfmask, set_sfmask);
328reg64!(kernel_gs_base, set_kernel_gs_base);
329reg64!(sysenter_cs, set_sysenter_cs);
330reg64!(sysenter_esp, set_sysenter_esp);
331reg64!(sysenter_eip, set_sysenter_eip);
332reg64!(cr2, set_cr2);
333reg64!(pat, set_pat);
334reg64!(spec_ctrl, set_spec_ctrl);
335reg32!(tsc_aux, set_tsc_aux);
336reg64!(rcx, set_rcx);
337reg64!(rdx, set_rdx);
338reg64!(rbx, set_rbx);
339reg64!(rbp, set_rbp);
340reg64!(rsi, set_rsi);
341reg64!(rdi, set_rdi);
342reg64!(r8, set_r8);
343reg64!(r9, set_r9);
344reg64!(r10, set_r10);
345reg64!(r11, set_r11);
346reg64!(r12, set_r12);
347reg64!(r13, set_r13);
348reg64!(r14, set_r14);
349reg64!(r15, set_r15);
350reg64!(next_rip, set_next_rip);
351reg64!(pcpu_id, set_pcpu_id);
352reg64!(xcr0, set_xcr0);
353
354#[cfg(test)]
355mod tests {
356 use super::*;
357
358 #[test]
359 fn test_reg_access() {
360 let nonce = 0xffff_ffff_ffff_ffffu64;
361 let nonce128 = ((nonce as u128) << 64) | nonce as u128;
362 let mut vmsa: SevVmsa = FromZeros::new_zeroed();
363 vmsa.register_protection_nonce = nonce;
364 let bitmap = [0xffu8; 64];
365 let mut vmsa_wrapper = VmsaWrapper {
366 vmsa: &mut vmsa,
367 bitmap: &bitmap,
368 };
369
370 let val = 0x0000_0055_0000_0055u128;
371 let val_xor = val ^ nonce128;
372 let cs = SevSelector::from(val);
373 let cs_xor = SevSelector::from(val_xor);
374 let vmpl = 2u8;
375 let rip = 0x55u64;
376 let rip_xor = rip ^ nonce;
377 let tsc = 0x55u32;
378 let tsc_xor = tsc ^ (nonce as u32);
379 let xmm_idx = 1;
380 let ymm_idx = 1;
381 let x87 = [0x55u64; 10];
382 let x87_xor = x87.map(|v| v ^ nonce);
383
384 vmsa_wrapper.set_cs(cs);
385 vmsa_wrapper.set_vmpl(vmpl);
386 vmsa_wrapper.set_rip(rip);
387 vmsa_wrapper.set_tsc_aux(tsc);
388 vmsa_wrapper.set_xmm_registers(xmm_idx, val);
389 vmsa_wrapper.set_ymm_registers(ymm_idx, val);
390 vmsa_wrapper.set_x87_registers(&x87);
391
392 assert!(vmsa_wrapper.cs() == cs);
393 assert!(vmsa_wrapper.vmpl() == vmpl);
394 assert!(vmsa_wrapper.rip() == rip);
395 assert!(vmsa_wrapper.xmm_registers(xmm_idx) == val);
396 assert!(vmsa_wrapper.ymm_registers(ymm_idx) == val);
397 assert!(vmsa_wrapper.tsc_aux() == tsc);
398 assert!(vmsa_wrapper.x87_registers() == x87);
399 assert!(vmsa.cs == cs_xor); assert!(vmsa.vmpl == vmpl); assert!(vmsa.rip == rip_xor); assert!(vmsa.tsc_aux == tsc_xor); assert!(vmsa.pkru == 0); assert!(vmsa.xmm_registers[xmm_idx].as_u128() == val_xor); assert!(vmsa.ymm_registers[ymm_idx].as_u128() == val_xor); assert!(vmsa.x87_registers == x87_xor);
407 }
408
409 #[test]
410 fn test_init() {
411 let mut vmsa: SevVmsa = FromZeros::new_zeroed();
412 let mut bitmap = [0x0u8; 64];
413 let xmm_idx = 1;
414 bitmap[5] = 0x80u8; bitmap[18] = 0x03u8; let mut vmsa_wrapper = VmsaWrapper {
417 vmsa: &mut vmsa,
418 bitmap: &bitmap,
419 };
420 vmsa_wrapper.reset(true);
421
422 assert!(vmsa_wrapper.rip() == 0);
423 assert!(vmsa_wrapper.xmm_registers(xmm_idx) == 0);
424
425 let nonce = vmsa.register_protection_nonce;
426 let xmm_val = ((nonce as u128) << 64) | nonce as u128;
427 assert!(vmsa.rip == nonce);
428 assert!(vmsa.xmm_registers[xmm_idx].as_u128() == xmm_val);
429 }
430}