Skip to main content

chia_datalayer_macro/
lib.rs

1use proc_macro::TokenStream;
2
3#[proc_macro_derive(PythonError)]
4pub fn python_error(input: TokenStream) -> TokenStream {
5    let input: syn::DeriveInput = syn::parse_macro_input!(input);
6    let mut output = TokenStream::new();
7
8    let syn::Data::Enum(input) = input.data else {
9        panic!("only enums are supported");
10    };
11
12    let names: Vec<proc_macro2::Ident> = input
13        .variants
14        .iter()
15        .map(|variant| quote::format_ident!("{}", variant.ident))
16        .collect();
17    let python_names: Vec<proc_macro2::Ident> = input
18        .variants
19        .iter()
20        .map(|variant| quote::format_ident!("{}Error", variant.ident))
21        .collect();
22
23    output.extend(TokenStream::from(quote::quote!(
24        #[cfg(feature = "py-bindings")]
25        pub mod python_exceptions {
26            use super::*;
27
28            #(
29                pyo3::create_exception!(chia_rs.datalayer, #python_names, pyo3::exceptions::PyException);
30            )*
31
32            pub fn add_to_module(py: pyo3::marker::Python<'_>, module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
33                use pyo3::prelude::PyModuleMethods;
34
35                #(
36                    module.add(stringify!(#python_names), py.get_type::<#python_names>())?;
37                )*
38
39                Ok(())
40            }
41        }
42
43        #[cfg(feature = "py-bindings")]
44        impl From<Error> for pyo3::PyErr {
45            fn from(err: Error) -> pyo3::PyErr {
46                let message = err.to_string();
47                match err {
48                    #(
49                        Error::#names(..) => python_exceptions::#python_names::new_err(message),
50                    )*
51                }
52            }
53        }
54    )));
55
56    output
57}