Skip to main content

somedb_macros/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::Field;
4
5#[proc_macro_derive(Storable)]
6pub fn derive_storable(item: TokenStream) -> TokenStream {
7    let input = syn::parse_macro_input!(item as syn::DeriveInput);
8
9    if input.generics.const_params().next().is_some()
10        || input.generics.lifetimes().next().is_some()
11        || input.generics.type_params().next().is_some()
12    {
13        panic!("Storables cannot have any generics.");
14    }
15
16    let ident = &input.ident;
17
18    match input.data {
19        syn::Data::Struct(s) => match s.fields {
20            syn::Fields::Named(n) => {
21                let names: Vec<_> = n.named.iter().map(|n| n.ident.as_ref().unwrap()).collect();
22                let types: Vec<_> = n.named.iter().map(|n| n.ty.clone()).collect();
23
24                quote! {
25                    #[automatically_derived]
26                    unsafe impl somedb::storable::Storable for #ident {
27                        fn type_hash() -> somedb::type_hash::TypeHash {
28                            use somedb::type_hash::TypeHash;
29                            let field_names = &[#(stringify!(#names)),*];
30                            let field_types = &[#(#types::type_hash()),*];
31
32                            unsafe {
33                                TypeHash::new(
34                                    std::any::type_name::<Self>(),
35                                    field_names,
36                                    field_types,
37                                )
38                            }
39                        }
40
41                        fn inner_encoded(&self) -> Vec<u8> {
42                            let mut bytes = Vec::new();
43                            #(bytes.append(&mut self.#names.encoded());)*
44
45                            bytes
46                        }
47
48                        fn decoded(mut reader: somedb::byte_reader::ByteReader) -> somedb::db::DbResult<Self> {
49                            #(let #names = #types::decoded(reader.reader_for_block())?;)*
50                            Ok(#ident {
51                                #(#names),*
52                            })
53                        }
54                    }
55                }
56                .into()
57            }
58            syn::Fields::Unit => panic!("Unit variants are not yet supported"),
59            syn::Fields::Unnamed(_) => panic!("Unnamed fields are not yet supported"),
60        },
61        syn::Data::Enum(_) => panic!("Enums are not yet supported."),
62        syn::Data::Union(_) => panic!("Unions are not yet supported."),
63    }
64}
65
66#[proc_macro_derive(Entity, attributes(entity_id))]
67pub fn derive_entity(item: TokenStream) -> TokenStream {
68    let input = syn::parse_macro_input!(item as syn::DeriveInput);
69
70    if input.generics.const_params().next().is_some()
71        || input.generics.lifetimes().next().is_some()
72        || input.generics.type_params().next().is_some()
73    {
74        panic!("Entities cannot have any generics (yet).");
75    }
76
77    let ident = &input.ident;
78
79    match input.data {
80        syn::Data::Struct(s) => match s.fields {
81            syn::Fields::Named(n) => {
82                let id_field: &Field = n
83                    .named
84                    .iter()
85                    .find(|f| {
86                        f.attrs
87                            .iter()
88                            .find(|a| a.meta.path().is_ident("entity_id"))
89                            .is_some()
90                    })
91                    .expect("there must be an Id");
92
93                let generate_id = if id_field
94                    .attrs
95                    .iter()
96                    .find(|a| a.meta.path().is_ident("entity_id"))
97                    .unwrap()
98                    .parse_nested_meta(|m| {
99                        if m.path.is_ident("auto_generate") {
100                            Ok(())
101                        } else {
102                            Err(m.error("invalid entity_id attribute"))
103                        }
104                    })
105                    .is_ok()
106                {
107                    quote! {const GENERATE_ID: bool = true}
108                } else {
109                    quote! {const GENERATE_ID: bool = false}
110                };
111                let id_field_name = id_field.ident.clone().unwrap();
112                let id_field_type = id_field.ty.clone();
113
114                quote! {
115                    #[automatically_derived]
116                    impl somedb::entity::Entity for #ident {
117                        type Id = #id_field_type;
118                        #generate_id;
119
120                        fn get_id(&self) -> #id_field_type {
121                            self.#id_field_name
122                        }
123
124                        fn set_id(&mut self, id: Self::Id) {
125                            self.#id_field_name = id;
126                        }
127                    }
128                }
129                .into()
130            }
131            syn::Fields::Unit => panic!("Unit variants are not yet supported"),
132            syn::Fields::Unnamed(_) => panic!("Unnamed fields are not yet supported"),
133        },
134        syn::Data::Enum(_) => panic!("Enums are not yet supported."),
135        syn::Data::Union(_) => panic!("Unions are not yet supported."),
136    }
137}
138
139#[proc_macro_attribute]
140pub fn entity(
141    _metadata: proc_macro::TokenStream,
142    input: proc_macro::TokenStream,
143) -> proc_macro::TokenStream {
144    let input: proc_macro2::TokenStream = input.into();
145    // FIXME: Clone should only be derived if it isn't already since
146    //        this is really annoying right now.
147    let output = quote! {
148        #[derive(Clone, somedb::Storable, somedb::Entity)]
149        #input
150    };
151    output.into()
152}