use core::fmt;
use bitflags::bitflags;
use crate::{asm::asm_td_call, TdAttributes};
#[derive(Debug, PartialEq)]
pub enum TdCallError {
TdxNoValidVeInfo,
TdxOperandInvalid,
TdxOperandBusy,
TdxPageAlreadyAccepted,
TdxPageSizeMismatch,
TdxMetadataFieldIdIncorrect,
TdxMetadataFieldNotWritable,
TdxMetadataFieldNotReadable,
TdxMetadataFieldValueNotValid,
TdxOpStateIncorrect,
TdxOperandAddrRangeError,
TdxPageMetadataIncorrect,
TdxServtdInfoHashMismatch,
TdxServtdNotBound,
TdxServtdUuidMismatch,
TdxTargetUuidMismatch,
TdxTargetUuidUpdated,
TdxTdFatal,
TdxTdKeysNotConfigured,
TdxTdcsNotAllocated,
Other,
}
#[derive(Debug)]
pub enum InitError {
TdxVendorIdMismatch,
TdxCpuLeafIdTooLow,
TdxGetVpInfoError(TdCallError),
}
#[derive(Debug)]
pub struct TdgVpInfo {
pub gpaw: Gpaw,
pub attributes: TdAttributes,
pub num_vcpus: u32,
pub max_vcpus: u32,
pub vcpu_index: u32,
pub sys_rd: u32,
}
#[repr(C)]
#[derive(Debug)]
pub struct TdgVeInfo {
pub exit_reason: u32,
pub exit_qualification: u64,
pub guest_linear_address: u64,
pub guest_physical_address: u64,
pub exit_instruction_length: u32,
pub exit_instruction_info: u32,
}
#[repr(C)]
#[derive(Debug)]
pub struct TdReport {
pub report_mac: ReportMac,
pub tee_tcb_info: [u8; 239],
pub reserved: [u8; 17],
pub tdinfo: TdInfo,
}
pub struct PageAttr {
gpa_mapping: u64,
gpa_attr: GpaAttrAll,
}
#[derive(Clone)]
pub enum InvdTranslations {
NoInvalidation,
InvdTlbAndEpxe,
InvdTlb,
InvdTlbExpGlobalTranslations,
}
pub fn get_tdinfo() -> Result<TdgVpInfo, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VpInfo as u64,
..Default::default()
};
td_call(&mut args)?;
Ok(TdgVpInfo {
gpaw: Gpaw::from(args.rcx),
attributes: TdAttributes::from_bits_truncate(args.rdx),
num_vcpus: args.r8 as u32,
max_vcpus: (args.r8 >> 32) as u32,
vcpu_index: args.r9 as u32,
sys_rd: args.r10 as u32,
})
}
pub fn get_veinfo() -> Result<TdgVeInfo, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VpVeinfoGet as u64,
..Default::default()
};
td_call(&mut args)?;
Ok(TdgVeInfo {
exit_reason: args.rcx as u32,
exit_qualification: args.rdx,
guest_linear_address: args.r8,
guest_physical_address: args.r9,
exit_instruction_length: args.r10 as u32,
exit_instruction_info: (args.r10 >> 32) as u32,
})
}
pub unsafe fn accept_page(sept_level: u64, gpa: u64) -> Result<(), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MemPageAccept as u64,
rcx: sept_level | gpa,
..Default::default()
};
td_call(&mut args)
}
pub fn release_private_page() {
todo!()
}
pub fn read_page_attr(gpa: &[u8]) -> Result<PageAttr, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MemPageAttrRd as u64,
rcx: gpa.as_ptr() as u64,
..Default::default()
};
td_call(&mut args)?;
Ok(PageAttr {
gpa_mapping: args.rcx,
gpa_attr: GpaAttrAll::from(args.rdx),
})
}
pub fn write_page_attr(page_attr: PageAttr, attr_flags: u64) -> Result<PageAttr, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MemPageAttrWr as u64,
rcx: page_attr.gpa_mapping,
rdx: u64::from(page_attr.gpa_attr),
r8: attr_flags,
..Default::default()
};
td_call(&mut args)?;
Ok(PageAttr {
gpa_mapping: args.rcx,
gpa_attr: GpaAttrAll::from(args.rdx),
})
}
pub fn get_report(report_gpa: u64, data_gpa: u64) -> Result<(), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MrReport as u64,
rcx: report_gpa,
rdx: data_gpa,
..Default::default()
};
td_call(&mut args)
}
pub fn extend_rtmr(extend_data_gpa: u64, reg_idx: u64) -> Result<(), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MrRtmrExtend as u64,
rcx: extend_data_gpa,
rdx: reg_idx,
..Default::default()
};
td_call(&mut args)
}
pub unsafe fn verify_report(report_mac_gpa: &[u8]) -> Result<(), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::MrVerifyreport as u64,
rcx: report_mac_gpa.as_ptr() as u64,
..Default::default()
};
td_call(&mut args)
}
pub fn read_td_metadata(field_identifier: u64) -> Result<u64, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VmRd as u64,
rdx: field_identifier,
..Default::default()
};
td_call(&mut args).map(|_| args.r8)
}
pub fn write_td_metadata(
field_identifier: u64,
data: u64,
write_mask: u64,
) -> Result<u64, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VmWr as u64,
rdx: field_identifier,
r8: data,
r9: write_mask,
..Default::default()
};
td_call(&mut args).map(|_| args.r8)
}
pub fn read_vcpu_metadata(field_identifier: u64) -> Result<u64, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VpRd as u64,
rdx: field_identifier,
..Default::default()
};
td_call(&mut args).map(|_| args.r8)
}
pub fn write_vcpu_metadata(
field_identifier: u64,
data: u64,
write_mask: u64,
) -> Result<u64, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VpWr as u64,
rdx: field_identifier,
r8: data,
r9: write_mask,
..Default::default()
};
td_call(&mut args).map(|_| args.r8)
}
pub fn read_sys_metadata(field_identifier: u64) -> Result<u64, TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::SysRd as u64,
rdx: field_identifier,
..Default::default()
};
td_call(&mut args).map(|_| args.r8)
}
pub fn read_sys_metadata_all() {
todo!()
}
pub fn read_sys_metadata_multiple() {
todo!()
}
pub fn read_vm_metadata_multiple() {
todo!()
}
pub fn write_vm_metadata_multiple() {
todo!()
}
pub fn read_vp_metadata_multiple() {
todo!()
}
pub fn write_vp_metadata_multiple() {
todo!()
}
pub fn set_cpuidve(cpuidve_flag: CpuidveFlag) -> Result<(), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::VpCpuidveSet as u64,
rcx: cpuidve_flag.bits(),
..Default::default()
};
td_call(&mut args)
}
pub fn assign_svn() {
todo!()
}
pub fn get_sealing_key() {}
pub fn read_servetd(
binding_handle: u64,
field_identifier: u64,
uuid: [u64; 4],
) -> Result<(u64, u64, [u64; 4]), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::ServetdRd as u64,
rcx: binding_handle,
rdx: field_identifier,
r10: uuid[0],
r11: uuid[1],
r12: uuid[2],
r13: uuid[3],
..Default::default()
};
td_call(&mut args).map(|_| (args.rdx, args.r8, [args.r10, args.r11, args.r12, args.r13]))
}
pub fn read_servetd_multiple() {
todo!()
}
pub fn write_servetd(
binding_handle: u64,
field_identifier: u64,
data: u64,
mask: u64,
uuid: [u64; 4],
) -> Result<(u64, [u64; 4]), TdCallError> {
let mut args = TdcallArgs {
rax: TdcallNum::ServetdWr as u64,
rcx: binding_handle,
rdx: field_identifier,
r8: data,
r9: mask,
r10: uuid[0],
r11: uuid[1],
r12: uuid[2],
r13: uuid[3],
};
td_call(&mut args).map(|_| (args.r8, [args.r10, args.r11, args.r12, args.r13]))
}
pub fn write_servetd_multiple() {
todo!()
}
pub fn enter_l2_vcpu(
l2_vm_idx: u64,
invd_translations: InvdTranslations,
guest_state_gpa: u64,
) -> Result<(), TdCallError> {
if (l2_vm_idx > 3) | (invd_translations.clone() as u64 > 3) {
return Err(TdCallError::TdxOperandInvalid);
}
let rcx = (l2_vm_idx << 52) | (invd_translations as u64);
let mut args = TdcallArgs {
rax: TdcallNum::VpEnter as u64,
rcx: l2_vm_idx,
rdx: guest_state_gpa,
..Default::default()
};
td_call(&mut args)
}
pub fn invalidate_l2_cached_ept(l2_vm_idx_bitmap: u64) -> Result<(), TdCallError> {
if l2_vm_idx_bitmap & !0b1110 == 0 {
return Err(TdCallError::TdxOperandInvalid);
}
let mut args = TdcallArgs {
rax: TdcallNum::VpInvept as u64,
rcx: l2_vm_idx_bitmap,
..Default::default()
};
td_call(&mut args)
}
pub fn invalidate_l2_gla(l2_vm_idx: u64, list: bool, gla: u64) -> Result<u64, TdCallError> {
if l2_vm_idx > 3 {
return Err(TdCallError::TdxOperandInvalid);
}
let rcx = (l2_vm_idx << 52) | if list { 0b1 } else { 0 };
let rdx = if list {
Gla::ListEntry(GlaListEntry(gla))
} else {
Gla::ListInfo(GlaListInfo(gla))
}
.value();
let mut args = TdcallArgs {
rax: TdcallNum::VpInvgla as u64,
rcx,
rdx,
..Default::default()
};
td_call(&mut args).map(|_| args.rdx)
}
fn td_call(args: &mut TdcallArgs) -> Result<(), TdCallError> {
let result = unsafe { asm_td_call(args) };
match result {
0 => Ok(()),
_ => Err((result >> 32).into()),
}
}
#[repr(u64)]
pub enum TdcallNum {
VpInfo = 1,
MrRtmrExtend = 2,
VpVeinfoGet = 3,
MrReport = 4,
VpCpuidveSet = 5,
MemPageAccept = 6,
VmRd = 7,
VmWr = 8,
VpRd = 9,
VpWr = 10,
SysRd = 11,
SysRdall = 12,
SysRdm = 13,
VmRdm = 14,
VmWrm = 15,
VpRdm = 16,
VpWrm = 17,
ServetdRd = 18,
ServetdRdm = 19,
ServetdWr = 20,
ServetdWrm = 21,
MrVerifyreport = 22,
MemPageAttrRd = 23,
MemPageAttrWr = 24,
VpEnter = 25,
VpInvept = 26,
VpInvgla = 27,
MrAssignsvns = 28,
MrKeyGet = 29,
MemPageRelease = 30,
}
#[repr(C)]
#[derive(Default)]
pub(crate) struct TdcallArgs {
rax: u64,
rcx: u64,
rdx: u64,
r8: u64,
r9: u64,
r10: u64,
r11: u64,
r12: u64,
r13: u64,
}
bitflags! {
pub struct CpuidveFlag: u64 {
const SUPERVISOR = 1 << 0;
const USER = 1 << 1;
}
}
bitflags! {
pub struct GpaAttr: u16 {
const R = 1;
const W = 1 << 1;
const XS = 1 << 2;
const XU = 1 << 3;
const VGP = 1 << 4;
const PWA = 1 << 5;
const SSS = 1 << 6;
const SVE = 1 << 7;
const VALID = 1 << 15;
}
}
pub struct GpaAttrAll {
l1_attr: GpaAttr,
vm1_attr: GpaAttr,
vm2_attr: GpaAttr,
vm3_attr: GpaAttr,
}
#[repr(C)]
#[derive(Debug)]
pub struct ReportMac {
pub report_type: ReportType,
pub cpu_svn: [u8; 16],
pub tee_tcb_info_hash: [u8; 48],
pub tee_info_hash: [u8; 48],
pub report_data: [u8; 64],
pub reserved: [u8; 32],
pub mac: [u8; 32],
}
#[repr(C)]
#[derive(Debug)]
pub struct ReportType {
pub tee_type: TeeType,
pub sub_type: u8,
pub version: u8,
pub reserved: u8,
}
#[derive(Debug)]
pub enum TeeType {
SGX,
TDX,
}
#[repr(C)]
#[derive(Debug)]
pub struct TdInfo {
pub attributes: u64,
pub xfam: u64,
pub mrtd: [u8; 48],
pub mr_config_id: [u8; 48],
pub mr_owner: [u8; 48],
pub mr_owner_config: [u8; 48],
pub rtmr0: [u8; 48],
pub rtmr1: [u8; 48],
pub rtmr2: [u8; 48],
pub rtmr3: [u8; 48],
pub servtd_hash: [u8; 48],
pub reserved: [u8; 64],
}
#[derive(Debug, Clone, Copy)]
pub enum Gpaw {
Bit48,
Bit52,
}
pub struct L2EnterGuestState {
pub rax: u64,
pub rcx: u64,
pub rdx: u64,
pub rbx: u64,
pub rsp: u64,
pub rbp: u64,
pub rsi: u64,
pub rdi: u64,
pub r8: u64,
pub r9: u64,
pub r10: u64,
pub r11: u64,
pub r12: u64,
pub r13: u64,
pub r14: u64,
pub r15: u64,
pub rflags: u64,
pub rip: u64,
pub ssp: u64,
pub guest_interrupt_status: u16,
}
enum Gla {
ListEntry(GlaListEntry),
ListInfo(GlaListInfo),
}
pub struct GlaListEntry(u64);
pub struct GlaListInfo(u64);
pub enum TdxVirtualExceptionType {
Hlt,
Io,
MsrRead,
MsrWrite,
CpuId,
VmCall,
Mwait,
Monitor,
EptViolation,
Wbinvd,
Rdpmc,
Other,
}
impl From<u64> for Gpaw {
fn from(val: u64) -> Self {
match val {
48 => Self::Bit48,
52 => Self::Bit52,
_ => panic!("Invalid gpaw"),
}
}
}
impl From<Gpaw> for u64 {
fn from(s: Gpaw) -> Self {
match s {
Gpaw::Bit48 => 48,
Gpaw::Bit52 => 52,
}
}
}
impl fmt::Display for Gpaw {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
Gpaw::Bit48 => write!(f, "48-bit"),
Gpaw::Bit52 => write!(f, "52-bit"),
}
}
}
impl From<u64> for GpaAttrAll {
fn from(val: u64) -> Self {
GpaAttrAll {
l1_attr: GpaAttr::from_bits_truncate((val & 0xFFFF) as u16),
vm1_attr: GpaAttr::from_bits_truncate(((val >> 16) & 0xFFFF) as u16),
vm2_attr: GpaAttr::from_bits_truncate(((val >> 32) & 0xFFFF) as u16),
vm3_attr: GpaAttr::from_bits_truncate(((val >> 48) & 0xFFFF) as u16),
}
}
}
impl From<GpaAttrAll> for u64 {
fn from(s: GpaAttrAll) -> Self {
let field1 = s.l1_attr.bits() as u64;
let field2 = (s.vm1_attr.bits() as u64) << 16;
let field3 = (s.vm2_attr.bits() as u64) << 32;
let field4 = (s.vm3_attr.bits() as u64) << 48;
field4 | field3 | field2 | field1
}
}
impl From<u32> for TdxVirtualExceptionType {
fn from(val: u32) -> Self {
match val {
10 => Self::CpuId,
12 => Self::Hlt,
15 => Self::Rdpmc,
18 => Self::VmCall,
30 => Self::Io,
31 => Self::MsrRead,
32 => Self::MsrWrite,
36 => Self::Mwait,
39 => Self::Monitor,
48 => Self::EptViolation,
54 => Self::Wbinvd,
_ => Self::Other,
}
}
}
impl From<TdCallError> for InitError {
fn from(error: TdCallError) -> Self {
InitError::TdxGetVpInfoError(error)
}
}
impl From<u64> for InvdTranslations {
fn from(val: u64) -> Self {
match val {
0 => Self::NoInvalidation,
1 => Self::InvdTlbAndEpxe,
2 => Self::InvdTlb,
3 => Self::InvdTlbExpGlobalTranslations,
_ => panic!("Invalid value"),
}
}
}
impl Gla {
fn value(&self) -> u64 {
match *self {
Gla::ListEntry(GlaListEntry(value)) => value,
Gla::ListInfo(GlaListInfo(value)) => value,
}
}
}
impl From<u64> for TdCallError {
fn from(val: u64) -> Self {
match val {
0x0000_0B0A => Self::TdxPageAlreadyAccepted,
0x8000_0200 => Self::TdxOperandBusy,
0x8000_0810 => Self::TdxTdKeysNotConfigured,
0xC000_0100 => Self::TdxOperandInvalid,
0xC000_0101 => Self::TdxOperandAddrRangeError,
0xC000_0300 => Self::TdxPageMetadataIncorrect,
0xC000_0606 => Self::TdxTdcsNotAllocated,
0xC000_0608 => Self::TdxOpStateIncorrect,
0xC000_0704 => Self::TdxNoValidVeInfo,
0xC000_0B0B => Self::TdxPageSizeMismatch,
0xC000_0C00 => Self::TdxMetadataFieldIdIncorrect,
0xC000_0C01 => Self::TdxMetadataFieldNotWritable,
0xC000_0C02 => Self::TdxMetadataFieldNotReadable,
0xC000_0C03 => Self::TdxMetadataFieldValueNotValid,
0xC000_0D03 => Self::TdxServtdInfoHashMismatch,
0xC000_0D04 => Self::TdxServtdUuidMismatch,
0xC000_0D05 => Self::TdxServtdNotBound,
0xC000_0D07 => Self::TdxTargetUuidMismatch,
0xC000_0D08 => Self::TdxTargetUuidUpdated,
0xE000_0604 => Self::TdxTdFatal,
_ => Self::Other,
}
}
}