Skip to main content

rqlite_rs_macros/
lib.rs

1#![warn(clippy::pedantic, clippy::all)]
2
3use proc_macro::TokenStream;
4use quote::quote;
5use syn::{parse_macro_input, DeriveInput, Type};
6
7mod field_type;
8
9#[proc_macro_derive(FromRow)]
10/// Derives the `FromRow` trait for a struct.
11///
12/// # Panics
13///
14/// This function will panic if the field name cannot be converted to a string.
15pub fn derive_from_row(input: TokenStream) -> TokenStream {
16    let input = parse_macro_input!(input as DeriveInput);
17
18    if let syn::Data::Struct(data) = &input.data {
19        if let syn::Fields::Named(fields) = &data.fields {
20            let field_vals = fields.named.iter().map(|field| {
21                let Some(field_name) = field.ident.as_ref() else {
22                    return syn::Error::new_spanned(
23                        field,
24                        "Expected named field"
25                    ).to_compile_error();
26                };
27                let field_name_string = field_name.to_string();
28
29                if let Type::Path(type_path) = &field.ty {
30                    match field_type::FieldType::from_type_path(type_path) {
31                        field_type::FieldType::Option => quote! {
32                            #field_name: row.get_opt(#field_name_string)?
33                        },
34                        field_type::FieldType::Blob => {
35                            #[cfg(feature = "fast-blob")]
36                            quote! {
37                                #field_name: rqlite_rs::decode::decode_blob(&row.get::<String>(#field_name_string)?)?
38                            }
39                            #[cfg(not(feature = "fast-blob"))]
40                            quote! {
41                                #field_name: row.get(#field_name_string)?
42                            }
43                        },
44                        field_type::FieldType::Normal => quote! {
45                            #field_name: row.get(#field_name_string)?
46                        },
47                    }
48                } else { quote! {
49                    #field_name: row.get(#field_name_string)?
50                } }
51            });
52
53            let struct_name = &input.ident;
54
55            return TokenStream::from(quote!(
56                impl rqlite_rs::FromRow for #struct_name {
57                    fn from_row(row: rqlite_rs::Row) -> Result<Self, rqlite_rs::IntoTypedError> {
58                        Ok(#struct_name {
59                            #(#field_vals),*
60                        })
61                    }
62                }
63            ));
64        }
65
66        if let syn::Fields::Unnamed(fields) = &data.fields {
67            let field_vals = fields.unnamed.iter().enumerate().map(|(index, field)| {
68                let index = syn::Index::from(index);
69
70                if let Type::Path(type_path) = &field.ty {
71                    match field_type::FieldType::from_type_path(type_path) {
72                        field_type::FieldType::Option => quote! {
73                            row.get_by_index_opt(#index)?
74                        },
75                        field_type::FieldType::Blob => {
76                            #[cfg(feature = "fast-blob")]
77                            quote! {
78                                rqlite_rs::decode::decode_blob(&row.get_by_index::<String>(#index)?)?
79                            }
80                            #[cfg(not(feature = "fast-blob"))]
81                            quote! {
82                                row.get_by_index(#index)?
83                            }
84                        },
85                        field_type::FieldType::Normal => quote! {
86                            row.get_by_index(#index)?
87                        },
88                    }
89                } else {
90                    quote! {
91                        row.get_by_index(#index)?
92                    }
93                }
94            });
95
96            let struct_name = &input.ident;
97
98            return TokenStream::from(quote!(
99                impl rqlite_rs::FromRow for #struct_name {
100                    fn from_row(row: rqlite_rs::Row) -> Result<Self, rqlite_rs::IntoTypedError> {
101                        Ok(#struct_name(
102                            #(#field_vals),*
103                        ))
104                    }
105                }
106            ));
107        }
108    }
109
110    TokenStream::from(
111        syn::Error::new(
112            input.ident.span(),
113            "Only structs with named fields are supported for `#[derive(FromRow)]`",
114        )
115        .to_compile_error(),
116    )
117}