use consortium_macros_helpers::{collect_field_types, path_type_arguments};
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::quote;
use syn::{Attribute, DeriveInput, GenericParam, Ident, LitInt, Meta, Type, parse_quote};
const DEFAULT_MAX_SIZE: &str = "1024";
pub fn derive_tee_param(item: TokenStream2) -> TokenStream2 {
let input = match syn::parse2::<DeriveInput>(item) {
Ok(input) => input,
Err(err) => return err.into_compile_error(),
};
derive_tee_param_impl(input).unwrap_or_else(|e| e.into_compile_error())
}
fn derive_tee_param_impl(input: DeriveInput) -> syn::Result<TokenStream2> {
let name = &input.ident;
let max_size = parse_tee_max_size(&input.attrs)?;
let field_types = collect_field_types(&input.data);
let mut combined: Option<syn::Error> = None;
let mut all_warnings: Vec<TokenStream2> = vec![];
for ty in field_types {
let (errs, warns) = check_tee_type(ty);
for err in errs {
match &mut combined {
None => combined = Some(err),
Some(prev) => prev.combine(err),
}
}
all_warnings.extend(warns);
}
if let Some(err) = combined {
return Err(err);
}
all_warnings.sort_by_key(|warning| warning.to_string());
all_warnings.dedup_by_key(|warning| warning.to_string());
let (_, orig_ty_generics, _) = input.generics.split_for_impl();
let codec_param = Ident::new("__TeeCodec", Span::call_site());
let mut generics = input.generics.clone();
for param in &mut generics.params {
if let GenericParam::Type(tp) = param {
tp.bounds
.push(parse_quote!(::consortium_tee::TeeParam<#codec_param>));
}
}
generics.params.push(parse_quote!(#codec_param));
let where_clause = generics.make_where_clause();
where_clause.predicates.push(parse_quote!(
#codec_param: for<'__buf> ::consortium_codec::CodecFor<
#name #orig_ty_generics,
Decoded<'__buf> = #name #orig_ty_generics,
>
));
where_clause
.predicates
.push(parse_quote!(#name #orig_ty_generics: 'static));
let (impl_generics, _, where_clause) = generics.split_for_impl();
Ok(quote! {
#(#all_warnings)*
impl #impl_generics ::consortium_tee::TeeParam<#codec_param>
for #name #orig_ty_generics
#where_clause
{
const MAX_SIZE: usize = #max_size;
fn into_tee_param(
&self,
) -> ::core::result::Result<::consortium_tee::TeeParamRepr, ::consortium_tee::TeeParamError>
{
let mut __buf = ::alloc::vec![0u8; <Self as ::consortium_tee::TeeParam<#codec_param>>::MAX_SIZE];
<#codec_param as ::consortium_codec::CodecFor<Self>>::encode(self, &mut __buf)
.map(|__len| {
__buf.truncate(__len);
::consortium_tee::TeeParamSlot::Serialized(__buf)
})
.map_err(|_| ::consortium_tee::TeeParamError::EncodeFailed)
}
fn from_tee_param(
repr: ::consortium_tee::TeeParamReprRef<'_>,
) -> ::core::result::Result<Self, ::consortium_tee::TeeParamError>
{
match repr {
::consortium_tee::TeeParamSlot::Serialized(__bytes) => {
<#codec_param as ::consortium_codec::CodecFor<Self>>::decode(__bytes)
.map_err(|_| ::consortium_tee::TeeParamError::DecodeFailed)
}
::consortium_tee::TeeParamSlot::Primitive { .. } => {
Err(::consortium_tee::TeeParamError::DecodeFailed)
}
}
}
}
})
}
fn parse_tee_max_size(attrs: &[Attribute]) -> syn::Result<LitInt> {
let mut max_size = None;
for attr in attrs {
if !attr.path().is_ident("tee") {
continue;
}
if matches!(attr.meta, Meta::Path(_)) {
continue;
}
attr.parse_nested_meta(|meta| {
if !meta.path.is_ident("max_size") {
return Err(meta.error("expected `max_size`"));
}
if max_size.is_some() {
return Err(meta.error("duplicate `max_size`"));
}
let value = meta.value()?;
let lit: LitInt = value.parse()?;
lit.base10_parse::<usize>()?;
max_size = Some(lit);
Ok(())
})?;
}
Ok(max_size.unwrap_or_else(|| LitInt::new(DEFAULT_MAX_SIZE, Span::call_site())))
}
fn check_tee_type(ty: &Type) -> (Vec<syn::Error>, Vec<TokenStream2>) {
let mut errors: Vec<syn::Error> = vec![];
let mut warnings: Vec<TokenStream2> = vec![];
match ty {
Type::Ptr(_) => errors.push(syn::Error::new_spanned(
ty,
"raw pointers are not allowed in TeeParam types: they cannot be serialized \
across the TEE secure-world boundary",
)),
Type::Reference(_) => errors.push(syn::Error::new_spanned(
ty,
"bare references are not allowed as TeeParam fields: use owned types instead",
)),
Type::BareFn(_) => errors.push(syn::Error::new_spanned(
ty,
"function pointers are not allowed in TeeParam types: they are not serializable",
)),
Type::Path(tp) => {
if let Some(ident) = tp.path.get_ident() {
match ident.to_string().as_str() {
"usize" | "isize" => {
errors.push(syn::Error::new_spanned(
ty,
format!(
"`{}` is pointer-width dependent and not allowed in TeeParam \
types: the TA may be 32-bit while the CA is 64-bit; \
use `u32`/`u64` or `i32`/`i64`",
ident
),
));
return (errors, warnings);
}
"f32" | "f64" => {
errors.push(syn::Error::new_spanned(
ty,
format!(
"TEE value slots are `uint32_t` pairs; `{}` encoding is \
ambiguous across TEE boundaries; use `u32` directly or \
implement `TeeParam` with explicit encoding",
ident
),
));
return (errors, warnings);
}
"u128" | "i128" => {
errors.push(syn::Error::new_spanned(
ty,
format!(
"`{}` does not fit in two `u32` TEE Value fields; \
wrap it in a newtype and implement `TeeParam` using a \
Memref (serialized) slot",
ident
),
));
return (errors, warnings);
}
_ => {}
}
}
for inner in path_type_arguments(ty) {
let (e, w) = check_tee_type(inner);
errors.extend(e);
warnings.extend(w);
}
}
Type::Array(arr) => {
let (e, w) = check_tee_type(&arr.elem);
errors.extend(e);
warnings.extend(w);
}
Type::Slice(sl) => {
let (e, w) = check_tee_type(&sl.elem);
errors.extend(e);
warnings.extend(w);
}
Type::Tuple(t) => {
for elem in &t.elems {
let (e, w) = check_tee_type(elem);
errors.extend(e);
warnings.extend(w);
}
}
Type::Paren(p) => {
let (e, w) = check_tee_type(&p.elem);
errors.extend(e);
warnings.extend(w);
}
Type::Group(g) => {
let (e, w) = check_tee_type(&g.elem);
errors.extend(e);
warnings.extend(w);
}
Type::TraitObject(_) | Type::ImplTrait(_) => errors.push(syn::Error::new_spanned(
ty,
"trait objects and `impl Trait` are not allowed in TeeParam types: use a \
concrete, sized type",
)),
_ => {}
}
(errors, warnings)
}