use darling::{ast::Data, Error, FromDeriveInput, FromField, ToTokens};
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::{parse_macro_input, DeriveInput, Result};
#[proc_macro_derive(FromRow, attributes(from_row))]
pub fn derive_from_row(input: TokenStream) -> TokenStream {
let derive_input = parse_macro_input!(input as DeriveInput);
match try_derive_from_row(&derive_input) {
Ok(result) => result,
Err(err) => err.write_errors().into(),
}
}
fn try_derive_from_row(input: &DeriveInput) -> std::result::Result<TokenStream, Error> {
let from_row_derive = DeriveFromRow::from_derive_input(input)?;
Ok(from_row_derive.generate()?)
}
#[derive(Debug, FromDeriveInput)]
#[darling(
attributes(from_row),
forward_attrs(allow, doc, cfg),
supports(struct_named)
)]
struct DeriveFromRow {
ident: syn::Ident,
generics: syn::Generics,
data: Data<(), FromRowField>,
}
impl DeriveFromRow {
fn validate(&self) -> Result<()> {
for field in self.fields() {
field.validate()?;
}
Ok(())
}
fn predicates(&self) -> Result<Vec<TokenStream2>> {
let mut predicates = Vec::new();
for field in self.fields() {
field.add_predicates(&mut predicates)?;
}
Ok(predicates)
}
fn fields(&self) -> &[FromRowField] {
match &self.data {
Data::Struct(fields) => &fields.fields,
_ => panic!("invalid shape"),
}
}
fn generate(self) -> Result<TokenStream> {
self.validate()?;
let ident = &self.ident;
let (impl_generics, ty_generics, where_clause) = self.generics.split_for_impl();
let original_predicates = where_clause.clone().map(|w| &w.predicates).into_iter();
let predicates = self.predicates()?;
let from_row_fields = self
.fields()
.iter()
.map(|f| f.generate_from_row())
.collect::<syn::Result<Vec<_>>>()?;
let try_from_row_fields = self
.fields()
.iter()
.map(|f| f.generate_try_from_row())
.collect::<syn::Result<Vec<_>>>()?;
Ok(quote! {
impl #impl_generics postgres_from_row::FromRow for #ident #ty_generics where #(#original_predicates),* #(#predicates),* {
fn from_row(row: &postgres_from_row::tokio_postgres::Row) -> Self {
Self {
#(#from_row_fields),*
}
}
fn try_from_row(row: &postgres_from_row::tokio_postgres::Row) -> std::result::Result<Self, postgres_from_row::tokio_postgres::Error> {
Ok(Self {
#(#try_from_row_fields),*
})
}
}
}
.into())
}
}
#[derive(Debug, FromField)]
#[darling(attributes(from_row), forward_attrs(allow, doc, cfg))]
struct FromRowField {
ident: Option<syn::Ident>,
ty: syn::Type,
#[darling(default)]
flatten: bool,
try_from: Option<String>,
from: Option<String>,
rename: Option<String>,
}
impl FromRowField {
fn validate(&self) -> Result<()> {
if self.from.is_some() && self.try_from.is_some() {
return Err(Error::custom(
r#"can't combine `#[from_row(from = "..")]` with `#[from_row(try_from = "..")]`"#,
)
.into());
}
if self.rename.is_some() && self.flatten {
return Err(Error::custom(
r#"can't combine `#[from_row(flatten)]` with `#[from_row(rename = "..")]`"#,
)
.into());
}
Ok(())
}
fn target_ty(&self) -> Result<TokenStream2> {
if let Some(from) = &self.from {
Ok(from.parse()?)
} else if let Some(try_from) = &self.try_from {
Ok(try_from.parse()?)
} else {
Ok(self.ty.to_token_stream())
}
}
fn column_name(&self) -> String {
self.rename
.as_ref()
.map(Clone::clone)
.unwrap_or_else(|| self.ident.as_ref().unwrap().to_string())
}
fn add_predicates(&self, predicates: &mut Vec<TokenStream2>) -> Result<()> {
let target_ty = &self.target_ty()?;
let ty = &self.ty;
predicates.push(if self.flatten {
quote! (#target_ty: postgres_from_row::FromRow)
} else {
quote! (#target_ty: for<'a> postgres_from_row::tokio_postgres::types::FromSql<'a>)
});
if self.from.is_some() {
predicates.push(quote!(#ty: std::convert::From<#target_ty>))
} else if self.try_from.is_some() {
let try_from = quote!(std::convert::TryFrom<#target_ty>);
predicates.push(quote!(#ty: #try_from));
predicates.push(quote!(postgres_from_row::tokio_postgres::Error: std::convert::From<<#ty as #try_from>::Error>));
predicates.push(quote!(<#ty as #try_from>::Error: std::fmt::Debug));
}
Ok(())
}
fn generate_from_row(&self) -> Result<TokenStream2> {
let ident = self.ident.as_ref().unwrap();
let column_name = self.column_name();
let field_ty = &self.ty;
let target_ty = self.target_ty()?;
let mut base = if self.flatten {
quote!(<#target_ty as postgres_from_row::FromRow>::from_row(row))
} else {
quote!(postgres_from_row::tokio_postgres::Row::get::<&str, #target_ty>(row, #column_name))
};
if self.from.is_some() {
base = quote!(<#field_ty as std::convert::From<#target_ty>>::from(#base));
} else if self.try_from.is_some() {
base = quote!(<#field_ty as std::convert::TryFrom<#target_ty>>::try_from(#base).expect("could not convert column"));
};
Ok(quote!(#ident: #base))
}
fn generate_try_from_row(&self) -> Result<TokenStream2> {
let ident = self.ident.as_ref().unwrap();
let column_name = self.column_name();
let field_ty = &self.ty;
let target_ty = self.target_ty()?;
let mut base = if self.flatten {
quote!(<#target_ty as postgres_from_row::FromRow>::try_from_row(row)?)
} else {
quote!(postgres_from_row::tokio_postgres::Row::try_get::<&str, #target_ty>(row, #column_name)?)
};
if self.from.is_some() {
base = quote!(<#field_ty as std::convert::From<#target_ty>>::from(#base));
} else if self.try_from.is_some() {
base = quote!(<#field_ty as std::convert::TryFrom<#target_ty>>::try_from(#base)?);
};
Ok(quote!(#ident: #base))
}
}