Skip to main content

hcl/
vmsa.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! Interface to `VmsaWrapper`, which combines a SEV-SNP VMSA
5//! with a bitmap to allow for register protection.
6
7use 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
22/// VMSA and register tweak bitmap.
23pub struct VmsaWrapper<'a, T> {
24    vmsa: T,
25    bitmap: &'a [u8; 64],
26}
27
28impl<'a, T> VmsaWrapper<'a, T> {
29    /// Create a VmsaWrapper
30    pub(crate) fn new(vmsa: T, bitmap: &'a [u8; 64]) -> Self {
31        VmsaWrapper { vmsa, bitmap }
32    }
33}
34
35/// Wraps a SEV VMSA structure with the register tweak bitmap to provide safe access methods.
36impl<T: Deref<Target = SevVmsa>> VmsaWrapper<'_, T> {
37    /// 64 bit register read
38    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    /// 32 bit register read
49    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    /// 128 bit register read
54    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    /// Gets an XMM VMSA register as u128
59    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    /// Gets a YMM VMSA register as u128
66    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    /// Gets the x87 VMSA registers
73    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
79/// Wraps a mutable SEV VMSA structure with the register tweak bitmap to provide safe access methods.
80impl<T: DerefMut<Target = SevVmsa>> VmsaWrapper<'_, T> {
81    /// 64 bit value to set in register
82    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    /// 32 bit value to set in register
91    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    /// 128 bit value to set in register
97    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    /// Create a new VMSA
103    pub fn reset(&mut self, vmsa_reg_prot: bool) {
104        *self.vmsa = FromZeros::new_zeroed();
105        if vmsa_reg_prot {
106            // Initialize nonce and all protected fields.
107            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                // Ensure direct accesses are not included in bitmap.
119                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    /// Sets an XMM VMSA register from u128
138    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    /// Sets an XMM VMSA register from u128
147    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    /// Sets the x87 registers
156    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    /// Atomically test and set the guest busy bit in v_intr_cntrl.
165    pub fn guest_busy_bit_test_and_set(&mut self) -> bool {
166        const VINTR_GUEST_BUSYBIT_MASK: u64 = 1u64 << 63;
167        // SAFETY: `v_intr_cntrl` is in the per-VP per-VTL VMSA. The `&mut self`
168        // guarantees no other Rust reference to this field exists, and this code
169        // only runs on the owning VP's thread. The atomic is needed because the
170        // untrusted hypervisor's hardware may concurrently access this field via
171        // VMRUN on another physical CPU.
172        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
178/// Check bitmap to see if a register is included in masking.
179fn 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            /// Gets a SevSelector VMSA register
189            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            /// Sets a SevSelector VMSA register
195            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            /// Gets a VMSA register
209            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            /// Sets a VMSA register
215            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            /// Gets a VMSA register
227            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            /// Sets a VMSA register
233            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            /// Gets a VMSA register directly
245            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            /// Sets a VMSA register directly
257            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            /// Access VMSA field directly in order to manipulate fields.
269            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); // bitmask applied to u128
400        assert!(vmsa.vmpl == vmpl); // no bitmask applied
401        assert!(vmsa.rip == rip_xor); // bitmask applied
402        assert!(vmsa.tsc_aux == tsc_xor); // bitmask applied to u32
403        assert!(vmsa.pkru == 0); // untouched
404        assert!(vmsa.xmm_registers[xmm_idx].as_u128() == val_xor); // bitmask applied to correct XMM offset
405        assert!(vmsa.ymm_registers[ymm_idx].as_u128() == val_xor); // bitmask applied to correct YMM offset
406        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; // rip
415        bitmap[18] = 0x03u8; // xmm_registers[1]
416        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}