use crate::FieldIr;
use proc_macro2::TokenStream;
use quote::quote;
use syn::{DeriveInput, Generics, Ident};
pub(crate) struct DeriveEncode {
ident: Ident,
generics: Generics,
fields: Vec<FieldIr>,
}
impl DeriveEncode {
pub fn new(input: DeriveInput) -> syn::Result<Self> {
let data = match input.data {
syn::Data::Struct(data) => data,
_ => abort!(
input.ident,
"can't derive `Encode` on this type: only `struct` types are allowed",
),
};
let fields = FieldIr::from_fields(data.fields)?;
Ok(Self {
ident: input.ident,
generics: input.generics.clone(),
fields,
})
}
pub fn to_tokens(&self) -> TokenStream {
let ident = &self.ident;
let (_, generics, where_clause) = self.generics.split_for_impl();
let mut lowerer = FieldLowerer::new();
for field in &self.fields {
lowerer.add_field(field);
}
let (encoded_len_body, encode_body) = lowerer.into_tokens();
quote! {
impl #generics ::ssh_encoding::Encode for #ident #generics #where_clause {
fn encoded_len(&self) -> ssh_encoding::Result<usize> {
[
#(#encoded_len_body)*,
]
.checked_sum()
}
fn encode(&self, writer: &mut impl Writer) -> ssh_encoding::Result<()> {
#(#encode_body)*;
Ok(())
}
}
}
}
}
struct FieldLowerer {
encoded_len_body: Vec<TokenStream>,
encode_body: Vec<TokenStream>,
}
impl FieldLowerer {
fn new() -> Self {
Self {
encoded_len_body: Vec::default(),
encode_body: Vec::default(),
}
}
fn add_field(&mut self, field: &FieldIr) {
let ident = field.ident.clone();
let field_length = quote! { self.#ident.encoded_len()? };
self.encoded_len_body.push(field_length);
let field_encoder = quote! { self.#ident.encode()? };
self.encode_body.push(field_encoder);
}
fn into_tokens(self) -> (Vec<TokenStream>, Vec<TokenStream>) {
(self.encoded_len_body, self.encode_body)
}
}