Skip to main content

virt/x86/
snp.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! SEV-SNP initial VMSA conversion.
5//!
6//! [`vmsa_from_initial_regs`] is used by the shared direct-boot loader for SNP
7//! backends, including KVM and MSHV. [`state_from_vmsa`] is currently used by
8//! KVM to apply the loader-built VMSA through KVM's register-based launch path;
9//! MSHV imports the VMSA page directly.
10
11use super::SegmentRegister;
12use super::TableRegister;
13use super::X86InitialRegs;
14use super::vp;
15use thiserror::Error;
16use x86defs::SegmentAttributes;
17use x86defs::snp::SevSelector;
18use x86defs::snp::SevVmsa;
19use zerocopy::FromZeros;
20
21const _: () = assert!(size_of::<SevVmsa>() <= hvdef::HV_PAGE_SIZE as usize);
22
23/// The subset of initial VP state represented by a direct-boot SNP VMSA.
24#[derive(Debug, Copy, Clone, PartialEq, Eq)]
25pub struct SnpVmsaState {
26    /// Architectural register state.
27    pub registers: vp::Registers,
28    /// Page attribute table state.
29    pub pat: u64,
30    /// Extended control register 0.
31    pub xcr0: u64,
32}
33
34/// An error parsing a direct-boot SNP VMSA.
35#[derive(Debug, Error)]
36pub enum SnpVmsaError {
37    /// The VMSA contains state outside the supported direct-boot contract.
38    #[error("unsupported SNP VMSA state")]
39    UnsupportedState,
40}
41
42/// Builds the direct-boot SNP VMSA corresponding to `initial`.
43pub fn vmsa_from_initial_regs(initial: &X86InitialRegs) -> SevVmsa {
44    vmsa_from_state(&SnpVmsaState {
45        registers: initial.registers,
46        pat: initial.pat.value,
47        xcr0: x86defs::xsave::XFEATURE_X87,
48    })
49}
50
51/// Parses and validates a direct-boot SNP VMSA.
52pub fn state_from_vmsa(vmsa: &SevVmsa) -> Result<SnpVmsaState, SnpVmsaError> {
53    let state = SnpVmsaState {
54        registers: vp::Registers {
55            rax: vmsa.rax,
56            rcx: vmsa.rcx,
57            rdx: vmsa.rdx,
58            rbx: vmsa.rbx,
59            rbp: vmsa.rbp,
60            rsp: vmsa.rsp,
61            rsi: vmsa.rsi,
62            rdi: vmsa.rdi,
63            r8: vmsa.r8,
64            r9: vmsa.r9,
65            r10: vmsa.r10,
66            r11: vmsa.r11,
67            r12: vmsa.r12,
68            r13: vmsa.r13,
69            r14: vmsa.r14,
70            r15: vmsa.r15,
71            rip: vmsa.rip,
72            rflags: vmsa.rflags,
73            cs: segment_from_vmsa(vmsa.cs),
74            ds: segment_from_vmsa(vmsa.ds),
75            es: segment_from_vmsa(vmsa.es),
76            fs: segment_from_vmsa(vmsa.fs),
77            gs: segment_from_vmsa(vmsa.gs),
78            ss: segment_from_vmsa(vmsa.ss),
79            tr: segment_from_vmsa(vmsa.tr),
80            ldtr: segment_from_vmsa(vmsa.ldtr),
81            gdtr: table_from_vmsa(vmsa.gdtr),
82            idtr: table_from_vmsa(vmsa.idtr),
83            cr0: vmsa.cr0,
84            cr2: vmsa.cr2,
85            cr3: vmsa.cr3,
86            cr4: vmsa.cr4,
87            cr8: 0,
88            efer: vmsa.efer & !x86defs::X64_EFER_SVME,
89        },
90        pat: vmsa.pat,
91        xcr0: vmsa.xcr0,
92    };
93
94    if &vmsa_from_state(&state) != vmsa {
95        return Err(SnpVmsaError::UnsupportedState);
96    }
97
98    Ok(state)
99}
100
101fn vmsa_from_state(state: &SnpVmsaState) -> SevVmsa {
102    let registers = &state.registers;
103    let mut vmsa = SevVmsa::new_zeroed();
104
105    vmsa.es = segment_to_vmsa(registers.es);
106    vmsa.cs = segment_to_vmsa(registers.cs);
107    vmsa.ss = segment_to_vmsa(registers.ss);
108    vmsa.ds = segment_to_vmsa(registers.ds);
109    vmsa.fs = segment_to_vmsa(registers.fs);
110    vmsa.gs = segment_to_vmsa(registers.gs);
111    vmsa.gdtr = table_to_vmsa(registers.gdtr);
112    vmsa.ldtr = segment_to_vmsa(registers.ldtr);
113    vmsa.idtr = table_to_vmsa(registers.idtr);
114    vmsa.tr = segment_to_vmsa(registers.tr);
115    vmsa.cpl = SegmentAttributes::from(registers.cs.attributes).descriptor_privilege_level();
116    vmsa.efer = registers.efer | x86defs::X64_EFER_SVME;
117    vmsa.cr4 = registers.cr4;
118    vmsa.cr3 = registers.cr3;
119    vmsa.cr0 = registers.cr0;
120    vmsa.rflags = registers.rflags;
121    vmsa.rip = registers.rip;
122    vmsa.rsp = registers.rsp;
123    vmsa.rax = registers.rax;
124    vmsa.cr2 = registers.cr2;
125    vmsa.pat = state.pat;
126    vmsa.rcx = registers.rcx;
127    vmsa.rdx = registers.rdx;
128    vmsa.rbx = registers.rbx;
129    vmsa.rbp = registers.rbp;
130    vmsa.rsi = registers.rsi;
131    vmsa.rdi = registers.rdi;
132    vmsa.r8 = registers.r8;
133    vmsa.r9 = registers.r9;
134    vmsa.r10 = registers.r10;
135    vmsa.r11 = registers.r11;
136    vmsa.r12 = registers.r12;
137    vmsa.r13 = registers.r13;
138    vmsa.r14 = registers.r14;
139    vmsa.r15 = registers.r15;
140    vmsa.sev_features.set_snp(true);
141    vmsa.xcr0 = state.xcr0;
142
143    vmsa
144}
145
146fn segment_to_vmsa(register: SegmentRegister) -> SevSelector {
147    SevSelector {
148        selector: register.selector,
149        attrib: (register.attributes & 0xff) | ((register.attributes >> 4) & 0xf00),
150        limit: register.limit,
151        base: register.base,
152    }
153}
154
155fn segment_from_vmsa(selector: SevSelector) -> SegmentRegister {
156    SegmentRegister {
157        selector: selector.selector,
158        attributes: (selector.attrib & 0xff) | ((selector.attrib & 0xf00) << 4),
159        limit: selector.limit,
160        base: selector.base,
161    }
162}
163
164fn table_to_vmsa(register: TableRegister) -> SevSelector {
165    SevSelector {
166        selector: 0,
167        attrib: 0,
168        limit: register.limit as u32,
169        base: register.base,
170    }
171}
172
173fn table_from_vmsa(selector: SevSelector) -> TableRegister {
174    TableRegister {
175        limit: selector.limit as u16,
176        base: selector.base,
177    }
178}
179
180#[cfg(test)]
181mod tests {
182    use super::*;
183
184    #[test]
185    fn direct_boot_vmsa_round_trips() {
186        let initial = X86InitialRegs {
187            registers: vp::Registers {
188                rax: 1,
189                rdx: 2,
190                rip: 0x100000,
191                rflags: 2,
192                cs: SegmentRegister {
193                    selector: 0x10,
194                    attributes: 0xa09b,
195                    limit: u32::MAX,
196                    base: 0,
197                },
198                efer: x86defs::X64_EFER_LME | x86defs::X64_EFER_LMA,
199                ..Default::default()
200            },
201            mtrrs: Default::default(),
202            pat: vp::Pat {
203                value: 0x7040600070406,
204            },
205        };
206
207        let vmsa = vmsa_from_initial_regs(&initial);
208        let state = state_from_vmsa(&vmsa).unwrap();
209
210        assert_eq!(state.registers, initial.registers);
211        assert_eq!(state.pat, initial.pat.value);
212        assert_eq!(state.xcr0, x86defs::xsave::XFEATURE_X87);
213        assert!(vmsa.sev_features.snp());
214        assert_ne!(vmsa.efer & x86defs::X64_EFER_SVME, 0);
215    }
216
217    #[test]
218    fn rejects_unsupported_state() {
219        let initial = X86InitialRegs {
220            registers: Default::default(),
221            mtrrs: Default::default(),
222            pat: Default::default(),
223        };
224        let mut vmsa = vmsa_from_initial_regs(&initial);
225        vmsa.virtual_tom = 0x1000;
226
227        assert!(matches!(
228            state_from_vmsa(&vmsa),
229            Err(SnpVmsaError::UnsupportedState)
230        ));
231    }
232}