extern crate proc_macro;
use proc_macro::TokenStream;
use quote::quote;
use syn::{Data, DeriveInput, Fields, parse_macro_input};
#[proc_macro_derive(FromRow, attributes(bigquery))]
pub fn derive_from_row(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
derive_from_row_impl(input).into()
}
fn derive_from_row_impl(input: DeriveInput) -> proc_macro2::TokenStream {
let name = input.ident;
let body = match input.data {
Data::Struct(data) => match data.fields {
Fields::Named(fields) if !fields.named.is_empty() => {
for f in &fields.named {
if let Err(err) = get_field_name(f) {
return err.to_compile_error();
}
}
let field_initializations = fields.named.iter().map(|f| {
let field_name = f.ident.as_ref().expect("named field must have identifier");
let db_column_name = get_field_name(f).expect("validated above");
quote! {
#field_name: row.take(#db_column_name)?,
}
});
quote! {
Self {
#( #field_initializations )*
}
}
}
Fields::Unnamed(fields) if !fields.unnamed.is_empty() => {
if let Err(err) = reject_bigquery_attrs(&fields.unnamed) {
return err.to_compile_error();
}
let field_initializations = (0..fields.unnamed.len()).map(|idx| {
quote! {
row.take(#idx)?,
}
});
quote! {
Self(
#( #field_initializations )*
)
}
}
_ => {
return syn::Error::new_spanned(
name,
"FromRow can only be derived for non-empty structs",
)
.to_compile_error();
}
},
_ => {
return syn::Error::new_spanned(
name,
"FromRow can only be derived for non-empty structs",
)
.to_compile_error();
}
};
let generics = add_trait_bounds(input.generics);
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
quote! {
impl #impl_generics std::convert::TryFrom<google_cloud_bigquery::query::Row> for #name #ty_generics #where_clause {
type Error = google_cloud_bigquery::error::RowError;
fn try_from(mut row: google_cloud_bigquery::query::Row) -> std::result::Result<Self, Self::Error> {
std::result::Result::Ok(#body)
}
}
}
}
#[proc_macro_derive(FromSql, attributes(bigquery))]
pub fn derive_from_sql(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
derive_from_sql_impl(input).into()
}
fn derive_from_sql_impl(input: DeriveInput) -> proc_macro2::TokenStream {
let name = input.ident;
let body = match input.data {
Data::Struct(data) => match data.fields {
Fields::Named(fields) if !fields.named.is_empty() => {
for f in &fields.named {
if let Err(err) = get_field_name(f) {
return err.to_compile_error();
}
}
let field_initializations = fields.named.iter().map(|f| {
let field_name = f.ident.as_ref().expect("named field must have identifier");
let db_column_name = get_field_name(f).expect("validated above");
quote! {
#field_name: value.take(#db_column_name)?,
}
});
quote! {
Self {
#( #field_initializations )*
}
}
}
Fields::Unnamed(fields) if !fields.unnamed.is_empty() => {
if let Err(err) = reject_bigquery_attrs(&fields.unnamed) {
return err.to_compile_error();
}
let field_initializations = (0..fields.unnamed.len()).map(|idx| {
quote! {
value.take(#idx)?,
}
});
quote! {
Self(
#( #field_initializations )*
)
}
}
_ => {
return syn::Error::new_spanned(
name,
"FromSql can only be derived for non-empty structs",
)
.to_compile_error();
}
},
_ => {
return syn::Error::new_spanned(
name,
"FromSql can only be derived for non-empty structs",
)
.to_compile_error();
}
};
let generics = add_trait_bounds(input.generics);
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
quote! {
impl #impl_generics google_cloud_bigquery::query::FromSql for #name #ty_generics #where_clause {
fn from_value(mut value: google_cloud_bigquery::query::SqlValue) -> std::result::Result<Self, google_cloud_bigquery::error::ConvertError> {
std::result::Result::Ok(#body)
}
}
}
}
fn reject_bigquery_attrs<'a>(fields: impl IntoIterator<Item = &'a syn::Field>) -> syn::Result<()> {
for field in fields {
if let Some(attr) = field.attrs.iter().find(|a| a.path().is_ident("bigquery")) {
return Err(syn::Error::new_spanned(
attr,
"bigquery attributes are not supported on tuple struct fields",
));
}
}
Ok(())
}
fn add_trait_bounds(mut generics: syn::Generics) -> syn::Generics {
for param in &mut generics.params {
if let syn::GenericParam::Type(type_param) = param {
type_param
.bounds
.push(syn::parse_quote!(google_cloud_bigquery::query::FromSql));
}
}
generics
}
fn get_field_name(field: &syn::Field) -> syn::Result<String> {
for attr in &field.attrs {
if attr.path().is_ident("bigquery") {
let mut renamed = None;
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("rename") {
let value = meta.value()?;
let lit: syn::LitStr = value.parse()?;
renamed = Some(lit.value());
Ok(())
} else {
Err(meta.error("unsupported bigquery attribute"))
}
})?;
if let Some(name) = renamed {
return Ok(name);
}
}
}
Ok(syn::ext::IdentExt::unraw(
field
.ident
.as_ref()
.expect("named field must have identifier"),
)
.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
use syn::parse_quote;
use test_case::test_case;
fn extract_first_field(input: DeriveInput) -> syn::Field {
match input.data {
Data::Struct(s) => match s.fields {
Fields::Named(n) => n.named.into_iter().next().unwrap(),
_ => unreachable!(),
},
_ => unreachable!(),
}
}
#[test]
fn test_invalid_bigquery_attribute_typo_errors() {
let make_input = || -> DeriveInput {
parse_quote! {
struct MyRow {
#[bigquery(renam = "custom_col")]
field: i64,
}
}
};
let field = extract_first_field(make_input());
let err = get_field_name(&field).unwrap_err();
assert!(
err.to_string().contains("unsupported bigquery attribute"),
"{err}"
);
let row_tokens = derive_from_row_impl(make_input()).to_string();
assert!(row_tokens.contains("unsupported bigquery attribute"));
let sql_tokens = derive_from_sql_impl(make_input()).to_string();
assert!(sql_tokens.contains("unsupported bigquery attribute"));
}
#[test]
fn test_invalid_bigquery_attribute_non_string_value_errors() {
let field = extract_first_field(parse_quote! {
struct MyRow {
#[bigquery(rename = 123)]
field: i64,
}
});
assert!(get_field_name(&field).is_err());
}
#[test_case("struct Empty {}"; "empty named struct")]
#[test_case("struct EmptyTuple();"; "empty tuple struct")]
#[test_case("struct Unit;"; "unit struct")]
fn test_rejects_empty_structs(def: &str) -> Result<(), syn::Error> {
let row_err = derive_from_row_impl(syn::parse_str(def)?).to_string();
assert!(
row_err.contains("FromRow can only be derived for non-empty structs"),
"unexpected expansion for {def}: {row_err}"
);
let sql_err = derive_from_sql_impl(syn::parse_str(def)?).to_string();
assert!(
sql_err.contains("FromSql can only be derived for non-empty structs"),
"unexpected expansion for {def}: {sql_err}"
);
Ok(())
}
#[test]
fn test_rejects_bigquery_attribute_on_tuple_struct_field() -> Result<(), syn::Error> {
let def = r#"struct TupleWithAttr(#[bigquery(rename = "custom")] i64);"#;
let row_err = derive_from_row_impl(syn::parse_str(def)?).to_string();
assert!(
row_err.contains("bigquery attributes are not supported on tuple struct fields"),
"unexpected expansion: {row_err}"
);
let sql_err = derive_from_sql_impl(syn::parse_str(def)?).to_string();
assert!(
sql_err.contains("bigquery attributes are not supported on tuple struct fields"),
"unexpected expansion: {sql_err}"
);
Ok(())
}
#[test]
fn test_rejects_non_structs() -> Result<(), syn::Error> {
let row_err = derive_from_row_impl(syn::parse_str("enum Foo {}")?).to_string();
assert!(
row_err.contains("FromRow can only be derived for non-empty structs"),
"unexpected expansion: {row_err}"
);
let sql_err = derive_from_sql_impl(syn::parse_str("enum Foo {}")?).to_string();
assert!(
sql_err.contains("FromSql can only be derived for non-empty structs"),
"unexpected expansion: {sql_err}"
);
Ok(())
}
#[test]
fn test_generics_expansion() -> Result<(), syn::Error> {
let input1: DeriveInput = syn::parse_str("struct Wrapper<T> { val: T }")?;
let row_tokens = derive_from_row_impl(input1).to_string();
assert!(
row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Wrapper < T >"),
"unexpected row expansion: {row_tokens}"
);
let input2: DeriveInput = syn::parse_str("struct Wrapper<T> { val: T }")?;
let sql_tokens = derive_from_sql_impl(input2).to_string();
assert!(
sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T >"),
"unexpected sql expansion: {sql_tokens}"
);
Ok(())
}
#[test]
fn test_generics_expansion_with_where_clause() -> Result<(), syn::Error> {
let input1: DeriveInput =
syn::parse_str("struct Wrapper<T> where T: std::fmt::Debug { val: T }")?;
let row_tokens = derive_from_row_impl(input1).to_string();
assert!(
row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Wrapper < T > where T : std :: fmt :: Debug"),
"unexpected row expansion: {row_tokens}"
);
let input2: DeriveInput =
syn::parse_str("struct Wrapper<T> where T: std::fmt::Debug { val: T }")?;
let sql_tokens = derive_from_sql_impl(input2).to_string();
assert!(
sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T > where T : std :: fmt :: Debug"),
"unexpected sql expansion: {sql_tokens}"
);
Ok(())
}
#[test]
fn test_generics_expansion_with_default_type_param() -> Result<(), syn::Error> {
let input1: DeriveInput = syn::parse_str("struct Wrapper<T = i64> { val: T }")?;
let row_tokens = derive_from_row_impl(input1).to_string();
assert!(
row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Wrapper < T >"),
"unexpected row expansion: {row_tokens}"
);
let input2: DeriveInput = syn::parse_str("struct Wrapper<T = i64> { val: T }")?;
let sql_tokens = derive_from_sql_impl(input2).to_string();
assert!(
sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T >"),
"unexpected sql expansion: {sql_tokens}"
);
Ok(())
}
#[test]
fn test_generics_expansion_multiple_params() -> Result<(), syn::Error> {
let input1: DeriveInput = syn::parse_str("struct Pair<A, B> { a: A, b: B }")?;
let row_tokens = derive_from_row_impl(input1).to_string();
assert!(
row_tokens.contains("impl < A : google_cloud_bigquery :: query :: FromSql , B : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Pair < A , B >"),
"unexpected row expansion: {row_tokens}"
);
let input2: DeriveInput = syn::parse_str("struct Pair<A, B> { a: A, b: B }")?;
let sql_tokens = derive_from_sql_impl(input2).to_string();
assert!(
sql_tokens.contains("impl < A : google_cloud_bigquery :: query :: FromSql , B : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Pair < A , B >"),
"unexpected sql expansion: {sql_tokens}"
);
Ok(())
}
#[test]
fn test_generics_tuple_struct_expansion() -> Result<(), syn::Error> {
let input1: DeriveInput = syn::parse_str("struct TupleWrapper<T, U>(T, U);")?;
let row_tokens = derive_from_row_impl(input1).to_string();
assert!(
row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql , U : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for TupleWrapper < T , U >"),
"unexpected row expansion: {row_tokens}"
);
let input2: DeriveInput = syn::parse_str("struct TupleWrapper<T, U>(T, U);")?;
let sql_tokens = derive_from_sql_impl(input2).to_string();
assert!(
sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql , U : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for TupleWrapper < T , U >"),
"unexpected sql expansion: {sql_tokens}"
);
Ok(())
}
}