extern crate alloc;
use alloc::fmt;
use core::fmt::Write;
use log::{Log, Metadata, Record};
use x86_64::registers::rflags::{self, RFlags};
use crate::asm::asm_td_vmcall;
pub struct TdxLogger;
pub static TDX_LOGGER: TdxLogger = TdxLogger;
#[derive(Debug, PartialEq)]
pub enum TdVmcallError {
TdxRetry,
TdxOperandInvalid,
TdxGpaInuse,
TdxAlignError,
Other,
}
#[repr(C)]
#[derive(Debug, Default)]
pub struct CpuIdInfo {
pub eax: usize,
pub ebx: usize,
pub ecx: usize,
pub edx: usize,
}
pub enum IoSize {
Size1 = 1,
Size2 = 2,
Size4 = 4,
Size8 = 8,
}
pub enum Direction {
In,
Out,
}
pub enum Operand {
Dx,
Immediate,
}
pub fn cpuid(eax: u32, ecx: u32) -> Result<CpuIdInfo, TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Cpuid as u64,
r12: eax as u64,
r13: ecx as u64,
..Default::default()
};
td_vmcall(&mut args)?;
Ok(CpuIdInfo {
eax: args.r12 as usize,
ebx: args.r13 as usize,
ecx: args.r14 as usize,
edx: args.r15 as usize,
})
}
pub fn hlt() {
let interrupt_blocked = !rflags::read().contains(RFlags::INTERRUPT_FLAG);
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Hlt as u64,
r12: interrupt_blocked as u64,
..Default::default()
};
let _ = td_vmcall(&mut args);
}
macro_rules! io_read {
($port:expr, $ty:ty) => {{
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Io as u64,
r12: core::mem::size_of::<$ty>() as u64,
r13: IO_READ,
r14: $port as u64,
..Default::default()
};
td_vmcall(&mut args)?;
Ok(args.r11 as u32)
}};
}
macro_rules! io_write {
($port:expr, $byte:expr, $size:expr) => {{
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Io as u64,
r12: core::mem::size_of_val(&$byte) as u64,
r13: IO_WRITE,
r14: $port as u64,
r15: $byte as u64,
..Default::default()
};
td_vmcall(&mut args)
}};
}
pub fn io_read(size: IoSize, port: u16) -> Result<u32, TdVmcallError> {
match size {
IoSize::Size1 => io_read!(port, u8),
IoSize::Size2 => io_read!(port, u16),
IoSize::Size4 => io_read!(port, u32),
_ => unreachable!(),
}
}
pub fn io_write(size: IoSize, port: u16, byte: u32) -> Result<(), TdVmcallError> {
match size {
IoSize::Size1 => io_write!(port, byte as u8, u8),
IoSize::Size2 => io_write!(port, byte as u16, u16),
IoSize::Size4 => io_write!(port, byte, u32),
_ => unreachable!(),
}
}
pub unsafe fn read_mmio(size: IoSize, mmio_gpa: u64) -> Result<u64, TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::RequestMmio as u64,
r12: size as u64,
r13: 0,
r14: mmio_gpa,
..Default::default()
};
td_vmcall(&mut args)?;
Ok(args.r11)
}
pub unsafe fn write_mmio(size: IoSize, mmio_gpa: u64, data: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::RequestMmio as u64,
r12: size as u64,
r13: 1,
r14: mmio_gpa,
r15: data,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn map_gpa(gpa: u64, size: u64) -> Result<(), (u64, TdVmcallError)> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Mapgpa as u64,
r12: gpa,
r13: size,
..Default::default()
};
td_vmcall(&mut args).map_err(|e| (args.r11, e))
}
pub unsafe fn rdmsr(index: u32) -> Result<u64, TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Rdmsr as u64,
r12: index as u64,
..Default::default()
};
td_vmcall(&mut args)?;
Ok(args.r11)
}
pub unsafe fn wrmsr(index: u32, value: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Wrmsr as u64,
r12: index as u64,
r13: value,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn perform_cache_operation(cache_operation: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Wbinvd as u64,
r12: cache_operation,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn get_quote(shared_gpa: u64, size: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::GetQuote as u64,
r12: shared_gpa,
r13: size,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn setup_event_notify_interrupt(interrupt_vector: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::SetupEventNotifyInterrupt as u64,
r12: interrupt_vector,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn get_tdvmcall_info(interrupt_vector: u64) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::GetTdVmcallInfo as u64,
r12: 0,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn get_td_service(
shared_gpa_input: u64,
shared_gpa_output: u64,
interrupt_vector: u64,
time_out: u64,
) -> Result<(), TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::Service as u64,
r12: shared_gpa_input,
r13: shared_gpa_output,
r14: interrupt_vector,
r15: time_out,
..Default::default()
};
td_vmcall(&mut args)
}
pub fn report_fatal_error_simple(message: &str) -> ! {
report_fatal_error(None, Some(message))
}
pub fn report_fatal_error_with_shared_memory(shared_gpa: u64) -> ! {
report_fatal_error(Some(shared_gpa), None)
}
pub fn report_fatal_error_full(shared_gpa: u64, brief_message: &str) -> ! {
report_fatal_error(Some(shared_gpa), Some(brief_message))
}
pub fn report_fatal_error(shared_gpa: Option<u64>, msg: Option<&str>) -> ! {
let mut r12_value: u64 = 0;
let mut r13_value: u64 = 0;
if let Some(gpa) = shared_gpa {
r12_value |= 1 << 63;
r13_value = gpa;
}
let mut message_regs = [0u64; 8];
let qemu_bit_indices = [14, 15, 3, 7, 6, 8, 9, 2];
let mut expose_mask: u64 = 0xfc00;
if let Some(message) = msg {
let mut buffer = [0u8; 64];
let bytes = message.as_bytes();
let len = bytes.len().min(63);
buffer[..len].copy_from_slice(&bytes[..len]);
for i in 0..8 {
let start = i * 8;
let mut chunk = [0u8; 8];
chunk.copy_from_slice(&buffer[start..start + 8]);
message_regs[i] = u64::from_le_bytes(chunk);
if message_regs[i] != 0 {
expose_mask |= 1 << qemu_bit_indices[i];
}
}
}
loop {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::ReportFatalError as u64,
r12: r12_value,
r13: r13_value,
rcx: expose_mask,
r14: message_regs[0],
r15: message_regs[1],
rbx: message_regs[2],
rdi: message_regs[3],
rsi: message_regs[4],
r8: message_regs[5],
r9: message_regs[6],
rdx: message_regs[7],
..Default::default()
};
let _ = td_vmcall(&mut args);
unsafe {
core::arch::asm!("pause");
}
}
}
pub fn pconfig(leaf: u64, r13: u64, r14: u64, r15: u64) -> Result<u64, TdVmcallError> {
let mut args = TdVmcallArgs {
r11: TdVmcallNum::PConfig as u64,
r12: leaf,
r13,
r14,
r15,
rcx: 0,
..Default::default()
};
td_vmcall(&mut args)?;
Ok(args.r11)
}
pub fn print(args: fmt::Arguments) {
Serial
.write_fmt(args)
.expect("Failed to write to serial port");
}
#[macro_export]
macro_rules! serial_print {
($fmt: literal $(, $($arg: tt)+)?) => {
$crate::tdvmcall::print(format_args!($fmt $(, $($arg)+)?));
}
}
#[macro_export]
macro_rules! serial_println {
($fmt: literal $(, $($arg: tt)+)?) => {
$crate::tdvmcall::print(format_args!(concat!($fmt, "\n") $(, $($arg)+)?))
}
}
#[repr(u64)]
pub enum TdVmcallNum {
Cpuid = 0x0000a,
Hlt = 0x0000c,
Io = 0x0001e,
Rdmsr = 0x0001f,
Wrmsr = 0x00020,
RequestMmio = 0x00030,
Wbinvd = 0x00036,
PConfig = 0x00041,
GetTdVmcallInfo = 0x10000,
Mapgpa = 0x10001,
GetQuote = 0x10002,
ReportFatalError = 0x10003,
SetupEventNotifyInterrupt = 0x10004,
Service = 0x10005,
}
#[repr(C)]
#[derive(Default)]
pub(crate) struct TdVmcallArgs {
r8: u64,
r9: u64,
r10: u64,
r11: u64,
r12: u64,
r13: u64,
r14: u64,
r15: u64,
rbx: u64,
rcx: u64,
rdi: u64,
rsi: u64,
rdx: u64,
}
const SERIAL_IO_PORT: u16 = 0x3F8;
const SERIAL_LINE_STS: u16 = 0x3FD;
const IO_READ: u64 = 0;
const IO_WRITE: u64 = 1;
fn td_vmcall(args: &mut TdVmcallArgs) -> Result<(), TdVmcallError> {
let result = unsafe { asm_td_vmcall(args) };
match result {
0 => Ok(()),
_ => Err(result.into()),
}
}
struct Serial;
impl Write for Serial {
fn write_str(&mut self, s: &str) -> fmt::Result {
for &byte in s.as_bytes() {
io_write!(SERIAL_IO_PORT, byte, u8).unwrap();
}
Ok(())
}
}
impl From<u64> for TdVmcallError {
fn from(val: u64) -> Self {
match val {
0x1 => Self::TdxRetry,
0x8000_0000_0000_0000 => Self::TdxOperandInvalid,
0x8000_0000_0000_0001 => Self::TdxGpaInuse,
0x8000_0000_0000_0002 => Self::TdxAlignError,
_ => Self::Other,
}
}
}
impl Log for TdxLogger {
fn enabled(&self, metadata: &Metadata) -> bool {
true
}
fn log(&self, record: &Record) {
let level = record.level();
let _ = write_serial_safe(format_args!("{}: {}\n", level, record.args()));
}
fn flush(&self) {}
}
fn write_serial_safe(args: core::fmt::Arguments) -> bool {
use core::fmt::Write;
struct SafeSerial;
impl Write for SafeSerial {
fn write_str(&mut self, s: &str) -> core::fmt::Result {
for &byte in s.as_bytes() {
if io_write!(SERIAL_IO_PORT, byte, u8).is_err() {
return Err(core::fmt::Error);
}
}
Ok(())
}
}
SafeSerial.write_fmt(args).is_ok()
}