Skip to main content

tamasfe_schemars_derive/
lib.rs

1#![forbid(unsafe_code)]
2
3#[macro_use]
4extern crate quote;
5#[macro_use]
6extern crate syn;
7extern crate proc_macro;
8
9mod ast;
10mod attr;
11mod metadata;
12mod schema_exprs;
13
14use ast::*;
15use proc_macro2::TokenStream;
16
17#[proc_macro_derive(JsonSchema, attributes(schemars, serde))]
18pub fn derive_json_schema_wrapper(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
19    let input = parse_macro_input!(input as syn::DeriveInput);
20    derive_json_schema(input).into()
21}
22
23fn derive_json_schema(mut input: syn::DeriveInput) -> TokenStream {
24    if let Err(e) = attr::process_serde_attrs(&mut input) {
25        return compile_error(&e);
26    }
27
28    let cont = match Container::from_ast(&input) {
29        Ok(c) => c,
30        Err(e) => return compile_error(&e),
31    };
32
33    let default_crate_name: syn::Path = parse_quote!(schemars);
34    let crate_name = cont
35        .attrs
36        .crate_name
37        .as_ref()
38        .unwrap_or(&default_crate_name);
39
40    let mut gen = cont.generics.clone();
41
42    add_trait_bounds(&crate_name, &mut gen);
43
44    let type_name = &cont.ident;
45    let (impl_generics, ty_generics, where_clause) = gen.split_for_impl();
46
47    if let Some(transparent_field) = cont.transparent_field() {
48        let (ty, type_def) = schema_exprs::type_for_schema(crate_name, transparent_field, 0);
49        return quote! {
50            #[automatically_derived]
51            impl #impl_generics #crate_name::JsonSchema for #type_name #ty_generics #where_clause {
52                #type_def
53
54                fn is_referenceable() -> bool {
55                    <#ty as #crate_name::JsonSchema>::is_referenceable()
56                }
57
58                fn schema_name() -> std::string::String {
59                    <#ty as #crate_name::JsonSchema>::schema_name()
60                }
61
62                fn json_schema(gen: &mut #crate_name::gen::SchemaGenerator) -> #crate_name::schema::Schema {
63                    <#ty as #crate_name::JsonSchema>::json_schema(gen)
64                }
65
66                fn json_schema_for_flatten(gen: &mut #crate_name::gen::SchemaGenerator) -> #crate_name::schema::Schema {
67                    <#ty as #crate_name::JsonSchema>::json_schema_for_flatten(gen)
68                }
69
70                fn add_schema_as_property(
71                    gen: &mut #crate_name::gen::SchemaGenerator,
72                    parent: &mut #crate_name::schema::SchemaObject,
73                    name: String,
74                    metadata: Option<#crate_name::schema::Metadata>,
75                    required: bool,
76                ) {
77                    <#ty as #crate_name::JsonSchema>::add_schema_as_property(gen, parent, name, metadata, required)
78                }
79            };
80        };
81    }
82
83    let mut schema_base_name = cont.name();
84    let schema_is_renamed = *type_name != schema_base_name;
85
86    if !schema_is_renamed {
87        if let Some(path) = cont.serde_attrs.remote() {
88            if let Some(segment) = path.segments.last() {
89                schema_base_name = segment.ident.to_string();
90            }
91        }
92    }
93
94    let type_params: Vec<_> = cont.generics.type_params().map(|ty| &ty.ident).collect();
95    let schema_name = if type_params.is_empty() {
96        quote! {
97            #schema_base_name.to_owned()
98        }
99    } else if schema_is_renamed {
100        let mut schema_name_fmt = schema_base_name;
101        for tp in &type_params {
102            schema_name_fmt.push_str(&format!("{{{}:.0}}", tp));
103        }
104        quote! {
105            format!(#schema_name_fmt #(,#type_params=#type_params::schema_name())*)
106        }
107    } else {
108        let mut schema_name_fmt = schema_base_name;
109        schema_name_fmt.push_str("_for_{}");
110        schema_name_fmt.push_str(&"_and_{}".repeat(type_params.len() - 1));
111        quote! {
112            format!(#schema_name_fmt #(,#type_params::schema_name())*)
113        }
114    };
115
116    let schema_expr = schema_exprs::expr_for_container(&cont);
117
118    quote! {
119        #[automatically_derived]
120        #[allow(unused_braces)]
121        impl #impl_generics #crate_name::JsonSchema for #type_name #ty_generics #where_clause {
122            fn schema_name() -> std::string::String {
123                #schema_name
124            }
125
126            fn json_schema(gen: &mut #crate_name::gen::SchemaGenerator) -> #crate_name::schema::Schema {
127                #schema_expr
128            }
129        };
130    }
131}
132
133fn add_trait_bounds(crate_name: &syn::Path, generics: &mut syn::Generics) {
134    for param in &mut generics.params {
135        if let syn::GenericParam::Type(ref mut type_param) = *param {
136            type_param
137                .bounds
138                .push(parse_quote!(#crate_name::JsonSchema));
139        }
140    }
141}
142
143fn compile_error<'a>(errors: impl IntoIterator<Item = &'a syn::Error>) -> TokenStream {
144    let compile_errors = errors.into_iter().map(syn::Error::to_compile_error);
145    quote! {
146        #(#compile_errors)*
147    }
148}