google_cloud_bigquery_derive/
lib.rs1extern crate proc_macro;
18
19use proc_macro::TokenStream;
20use quote::quote;
21use syn::{Data, DeriveInput, Fields, parse_macro_input};
22
23#[proc_macro_derive(FromRow, attributes(bigquery))]
30pub fn derive_from_row(input: TokenStream) -> TokenStream {
31 let input = parse_macro_input!(input as DeriveInput);
32 derive_from_row_impl(input).into()
33}
34
35fn derive_from_row_impl(input: DeriveInput) -> proc_macro2::TokenStream {
36 let name = input.ident;
37
38 let body = match input.data {
39 Data::Struct(data) => match data.fields {
40 Fields::Named(fields) if !fields.named.is_empty() => {
41 for f in &fields.named {
42 if let Err(err) = get_field_name(f) {
43 return err.to_compile_error();
44 }
45 }
46 let field_initializations = fields.named.iter().map(|f| {
47 let field_name = f.ident.as_ref().expect("named field must have identifier");
48 let db_column_name = get_field_name(f).expect("validated above");
49 quote! {
50 #field_name: row.take(#db_column_name)?,
51 }
52 });
53 quote! {
54 Self {
55 #( #field_initializations )*
56 }
57 }
58 }
59 Fields::Unnamed(fields) if !fields.unnamed.is_empty() => {
60 if let Err(err) = reject_bigquery_attrs(&fields.unnamed) {
61 return err.to_compile_error();
62 }
63 let field_initializations = (0..fields.unnamed.len()).map(|idx| {
64 quote! {
65 row.take(#idx)?,
66 }
67 });
68 quote! {
69 Self(
70 #( #field_initializations )*
71 )
72 }
73 }
74 _ => {
75 return syn::Error::new_spanned(
76 name,
77 "FromRow can only be derived for non-empty structs",
78 )
79 .to_compile_error();
80 }
81 },
82 _ => {
83 return syn::Error::new_spanned(
84 name,
85 "FromRow can only be derived for non-empty structs",
86 )
87 .to_compile_error();
88 }
89 };
90
91 let generics = add_trait_bounds(input.generics);
94 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
95
96 quote! {
97 impl #impl_generics std::convert::TryFrom<google_cloud_bigquery::query::Row> for #name #ty_generics #where_clause {
98 type Error = google_cloud_bigquery::error::RowError;
99
100 fn try_from(mut row: google_cloud_bigquery::query::Row) -> std::result::Result<Self, Self::Error> {
101 std::result::Result::Ok(#body)
102 }
103 }
104 }
105}
106
107#[proc_macro_derive(FromSql, attributes(bigquery))]
113pub fn derive_from_sql(input: TokenStream) -> TokenStream {
114 let input = parse_macro_input!(input as DeriveInput);
115 derive_from_sql_impl(input).into()
116}
117
118fn derive_from_sql_impl(input: DeriveInput) -> proc_macro2::TokenStream {
119 let name = input.ident;
120
121 let body = match input.data {
122 Data::Struct(data) => match data.fields {
123 Fields::Named(fields) if !fields.named.is_empty() => {
124 for f in &fields.named {
125 if let Err(err) = get_field_name(f) {
126 return err.to_compile_error();
127 }
128 }
129 let field_initializations = fields.named.iter().map(|f| {
130 let field_name = f.ident.as_ref().expect("named field must have identifier");
131 let db_column_name = get_field_name(f).expect("validated above");
132 quote! {
133 #field_name: value.take(#db_column_name)?,
134 }
135 });
136 quote! {
137 Self {
138 #( #field_initializations )*
139 }
140 }
141 }
142 Fields::Unnamed(fields) if !fields.unnamed.is_empty() => {
143 if let Err(err) = reject_bigquery_attrs(&fields.unnamed) {
144 return err.to_compile_error();
145 }
146 let field_initializations = (0..fields.unnamed.len()).map(|idx| {
147 quote! {
148 value.take(#idx)?,
149 }
150 });
151 quote! {
152 Self(
153 #( #field_initializations )*
154 )
155 }
156 }
157 _ => {
158 return syn::Error::new_spanned(
159 name,
160 "FromSql can only be derived for non-empty structs",
161 )
162 .to_compile_error();
163 }
164 },
165 _ => {
166 return syn::Error::new_spanned(
167 name,
168 "FromSql can only be derived for non-empty structs",
169 )
170 .to_compile_error();
171 }
172 };
173
174 let generics = add_trait_bounds(input.generics);
175 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
176
177 quote! {
178 impl #impl_generics google_cloud_bigquery::query::FromSql for #name #ty_generics #where_clause {
179 fn from_value(mut value: google_cloud_bigquery::query::SqlValue) -> std::result::Result<Self, google_cloud_bigquery::error::ConvertError> {
180 std::result::Result::Ok(#body)
181 }
182 }
183 }
184}
185
186fn reject_bigquery_attrs<'a>(fields: impl IntoIterator<Item = &'a syn::Field>) -> syn::Result<()> {
187 for field in fields {
188 if let Some(attr) = field.attrs.iter().find(|a| a.path().is_ident("bigquery")) {
189 return Err(syn::Error::new_spanned(
190 attr,
191 "bigquery attributes are not supported on tuple struct fields",
192 ));
193 }
194 }
195 Ok(())
196}
197
198fn add_trait_bounds(mut generics: syn::Generics) -> syn::Generics {
200 for param in &mut generics.params {
201 if let syn::GenericParam::Type(type_param) = param {
202 type_param
203 .bounds
204 .push(syn::parse_quote!(google_cloud_bigquery::query::FromSql));
205 }
206 }
207 generics
208}
209
210fn get_field_name(field: &syn::Field) -> syn::Result<String> {
211 for attr in &field.attrs {
212 if attr.path().is_ident("bigquery") {
213 let mut renamed = None;
214 attr.parse_nested_meta(|meta| {
215 if meta.path.is_ident("rename") {
216 let value = meta.value()?;
217 let lit: syn::LitStr = value.parse()?;
218 renamed = Some(lit.value());
219 Ok(())
220 } else {
221 Err(meta.error("unsupported bigquery attribute"))
222 }
223 })?;
224 if let Some(name) = renamed {
225 return Ok(name);
226 }
227 }
228 }
229 Ok(syn::ext::IdentExt::unraw(
230 field
231 .ident
232 .as_ref()
233 .expect("named field must have identifier"),
234 )
235 .to_string())
236}
237
238#[cfg(test)]
239mod tests {
240 use super::*;
241 use syn::parse_quote;
242 use test_case::test_case;
243
244 fn extract_first_field(input: DeriveInput) -> syn::Field {
245 match input.data {
246 Data::Struct(s) => match s.fields {
247 Fields::Named(n) => n.named.into_iter().next().unwrap(),
248 _ => unreachable!(),
249 },
250 _ => unreachable!(),
251 }
252 }
253
254 #[test]
255 fn test_invalid_bigquery_attribute_typo_errors() {
256 let make_input = || -> DeriveInput {
257 parse_quote! {
258 struct MyRow {
259 #[bigquery(renam = "custom_col")]
260 field: i64,
261 }
262 }
263 };
264 let field = extract_first_field(make_input());
265
266 let err = get_field_name(&field).unwrap_err();
267 assert!(
268 err.to_string().contains("unsupported bigquery attribute"),
269 "{err}"
270 );
271
272 let row_tokens = derive_from_row_impl(make_input()).to_string();
273 assert!(row_tokens.contains("unsupported bigquery attribute"));
274
275 let sql_tokens = derive_from_sql_impl(make_input()).to_string();
276 assert!(sql_tokens.contains("unsupported bigquery attribute"));
277 }
278
279 #[test]
280 fn test_invalid_bigquery_attribute_non_string_value_errors() {
281 let field = extract_first_field(parse_quote! {
282 struct MyRow {
283 #[bigquery(rename = 123)]
284 field: i64,
285 }
286 });
287
288 assert!(get_field_name(&field).is_err());
289 }
290
291 #[test_case("struct Empty {}"; "empty named struct")]
292 #[test_case("struct EmptyTuple();"; "empty tuple struct")]
293 #[test_case("struct Unit;"; "unit struct")]
294 fn test_rejects_empty_structs(def: &str) -> Result<(), syn::Error> {
295 let row_err = derive_from_row_impl(syn::parse_str(def)?).to_string();
296 assert!(
297 row_err.contains("FromRow can only be derived for non-empty structs"),
298 "unexpected expansion for {def}: {row_err}"
299 );
300
301 let sql_err = derive_from_sql_impl(syn::parse_str(def)?).to_string();
302 assert!(
303 sql_err.contains("FromSql can only be derived for non-empty structs"),
304 "unexpected expansion for {def}: {sql_err}"
305 );
306 Ok(())
307 }
308
309 #[test]
310 fn test_rejects_bigquery_attribute_on_tuple_struct_field() -> Result<(), syn::Error> {
311 let def = r#"struct TupleWithAttr(#[bigquery(rename = "custom")] i64);"#;
312 let row_err = derive_from_row_impl(syn::parse_str(def)?).to_string();
313 assert!(
314 row_err.contains("bigquery attributes are not supported on tuple struct fields"),
315 "unexpected expansion: {row_err}"
316 );
317
318 let sql_err = derive_from_sql_impl(syn::parse_str(def)?).to_string();
319 assert!(
320 sql_err.contains("bigquery attributes are not supported on tuple struct fields"),
321 "unexpected expansion: {sql_err}"
322 );
323 Ok(())
324 }
325
326 #[test]
327 fn test_rejects_non_structs() -> Result<(), syn::Error> {
328 let row_err = derive_from_row_impl(syn::parse_str("enum Foo {}")?).to_string();
329 assert!(
330 row_err.contains("FromRow can only be derived for non-empty structs"),
331 "unexpected expansion: {row_err}"
332 );
333
334 let sql_err = derive_from_sql_impl(syn::parse_str("enum Foo {}")?).to_string();
335 assert!(
336 sql_err.contains("FromSql can only be derived for non-empty structs"),
337 "unexpected expansion: {sql_err}"
338 );
339 Ok(())
340 }
341
342 #[test]
343 fn test_generics_expansion() -> Result<(), syn::Error> {
344 let input1: DeriveInput = syn::parse_str("struct Wrapper<T> { val: T }")?;
345 let row_tokens = derive_from_row_impl(input1).to_string();
346 assert!(
347 row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Wrapper < T >"),
348 "unexpected row expansion: {row_tokens}"
349 );
350
351 let input2: DeriveInput = syn::parse_str("struct Wrapper<T> { val: T }")?;
352 let sql_tokens = derive_from_sql_impl(input2).to_string();
353 assert!(
354 sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T >"),
355 "unexpected sql expansion: {sql_tokens}"
356 );
357 Ok(())
358 }
359
360 #[test]
361 fn test_generics_expansion_with_where_clause() -> Result<(), syn::Error> {
362 let input1: DeriveInput =
363 syn::parse_str("struct Wrapper<T> where T: std::fmt::Debug { val: T }")?;
364 let row_tokens = derive_from_row_impl(input1).to_string();
365 assert!(
366 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"),
367 "unexpected row expansion: {row_tokens}"
368 );
369
370 let input2: DeriveInput =
371 syn::parse_str("struct Wrapper<T> where T: std::fmt::Debug { val: T }")?;
372 let sql_tokens = derive_from_sql_impl(input2).to_string();
373 assert!(
374 sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T > where T : std :: fmt :: Debug"),
375 "unexpected sql expansion: {sql_tokens}"
376 );
377 Ok(())
378 }
379
380 #[test]
381 fn test_generics_expansion_with_default_type_param() -> Result<(), syn::Error> {
382 let input1: DeriveInput = syn::parse_str("struct Wrapper<T = i64> { val: T }")?;
383 let row_tokens = derive_from_row_impl(input1).to_string();
384 assert!(
385 row_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > std :: convert :: TryFrom < google_cloud_bigquery :: query :: Row > for Wrapper < T >"),
386 "unexpected row expansion: {row_tokens}"
387 );
388
389 let input2: DeriveInput = syn::parse_str("struct Wrapper<T = i64> { val: T }")?;
390 let sql_tokens = derive_from_sql_impl(input2).to_string();
391 assert!(
392 sql_tokens.contains("impl < T : google_cloud_bigquery :: query :: FromSql > google_cloud_bigquery :: query :: FromSql for Wrapper < T >"),
393 "unexpected sql expansion: {sql_tokens}"
394 );
395 Ok(())
396 }
397
398 #[test]
399 fn test_generics_expansion_multiple_params() -> Result<(), syn::Error> {
400 let input1: DeriveInput = syn::parse_str("struct Pair<A, B> { a: A, b: B }")?;
401 let row_tokens = derive_from_row_impl(input1).to_string();
402 assert!(
403 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 >"),
404 "unexpected row expansion: {row_tokens}"
405 );
406
407 let input2: DeriveInput = syn::parse_str("struct Pair<A, B> { a: A, b: B }")?;
408 let sql_tokens = derive_from_sql_impl(input2).to_string();
409 assert!(
410 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 >"),
411 "unexpected sql expansion: {sql_tokens}"
412 );
413 Ok(())
414 }
415
416 #[test]
417 fn test_generics_tuple_struct_expansion() -> Result<(), syn::Error> {
418 let input1: DeriveInput = syn::parse_str("struct TupleWrapper<T, U>(T, U);")?;
419 let row_tokens = derive_from_row_impl(input1).to_string();
420 assert!(
421 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 >"),
422 "unexpected row expansion: {row_tokens}"
423 );
424
425 let input2: DeriveInput = syn::parse_str("struct TupleWrapper<T, U>(T, U);")?;
426 let sql_tokens = derive_from_sql_impl(input2).to_string();
427 assert!(
428 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 >"),
429 "unexpected sql expansion: {sql_tokens}"
430 );
431 Ok(())
432 }
433}