1use 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#[derive(Debug, Copy, Clone, PartialEq, Eq)]
25pub struct SnpVmsaState {
26 pub registers: vp::Registers,
28 pub pat: u64,
30 pub xcr0: u64,
32}
33
34#[derive(Debug, Error)]
36pub enum SnpVmsaError {
37 #[error("unsupported SNP VMSA state")]
39 UnsupportedState,
40}
41
42pub 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
51pub 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}