Skip to main content

rs_macros/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{parse_macro_input, Attribute, Data, DeriveInput, Fields, Meta};
4
5#[proc_macro_derive(JsonSchemaOneOf)]
6pub fn derive_json_schema_one_of(input: TokenStream) -> TokenStream {
7    let input = parse_macro_input!(input as DeriveInput);
8
9    let name = &input.ident;
10    let enum_name_str = name.to_string();
11
12    // Parse the enum data
13    let data = match &input.data {
14        Data::Enum(data) => data,
15        _ => {
16            return syn::Error::new_spanned(
17                &input.ident,
18                "JsonSchemaOneOf can only be derived for enums",
19            )
20            .to_compile_error()
21            .into();
22        }
23    };
24
25    // Check for unit variants only (no fields)
26    for variant in &data.variants {
27        match &variant.fields {
28            Fields::Unit => {}
29            _ => {
30                return syn::Error::new_spanned(
31                    &variant.ident,
32                    "JsonSchemaOneOf only supports unit variants (no fields)",
33                )
34                .to_compile_error()
35                .into();
36            }
37        }
38    }
39
40    // Extract serde rename_all if present
41    let rename_case = extract_serde_rename_all(&input.attrs); // Generate const values for each variant
42    let const_variants = data.variants.iter().map(|variant| {
43        let variant_name = variant.ident.to_string();
44        let json_name = apply_rename_case(&variant_name, &rename_case);
45
46        quote! {
47            schemars::schema::Schema::Object(schemars::schema::SchemaObject {
48                const_value: Some(serde_json::Value::String(#json_name.to_string())),
49                metadata: Some(Box::new(schemars::schema::Metadata {
50                    title: Some(#variant_name.to_string()),
51                    ..Default::default()
52                })),
53                ..Default::default()
54            })
55        }
56    });
57
58    let expanded = quote! {
59        impl schemars::JsonSchema for #name {
60            fn schema_name() -> String {
61                #enum_name_str.to_owned()
62            }
63
64            fn json_schema(_gen: &mut schemars::gen::SchemaGenerator) -> schemars::schema::Schema {
65                use schemars::schema::*;
66
67                let mut schema = SchemaObject::default();
68                schema.subschemas().one_of = Some(vec![
69                    #(#const_variants,)*
70                ]);
71
72                schemars::schema::Schema::Object(schema)
73            }
74        }
75    };
76
77    TokenStream::from(expanded)
78}
79
80/// Extract the serde rename_all case style from attributes
81fn extract_serde_rename_all(attrs: &[Attribute]) -> Option<String> {
82    for attr in attrs {
83        if attr.path().is_ident("serde") {
84            if let Meta::List(meta_list) = &attr.meta {
85                // Convert TokenStream to string and parse it
86                let tokens_str = meta_list.tokens.to_string();
87                if tokens_str.contains("rename_all") {
88                    // Parse rename_all = "camelCase"
89                    if let Some(start) = tokens_str.find('"') {
90                        if let Some(end) = tokens_str.rfind('"') {
91                            if start < end {
92                                return Some(tokens_str[start + 1..end].to_string());
93                            }
94                        }
95                    }
96                }
97            }
98        }
99    }
100    None
101}
102
103/// Apply serde rename case transformation to a variant name
104fn apply_rename_case(name: &str, rename_case: &Option<String>) -> String {
105    match rename_case.as_deref() {
106        Some("camelCase") => to_camel_case(name),
107        Some("snake_case") => to_snake_case(name),
108        Some("kebab-case") => to_kebab_case(name),
109        Some("SCREAMING_SNAKE_CASE") => to_screaming_snake_case(name),
110        Some("PascalCase") => to_pascal_case(name),
111        Some("SCREAMING-KEBAB-CASE") => to_screaming_kebab_case(name),
112        _ => name.to_string(),
113    }
114}
115
116/// Convert PascalCase to camelCase
117fn to_camel_case(s: &str) -> String {
118    if s.is_empty() {
119        return s.to_string();
120    }
121
122    let mut result = String::new();
123    let mut chars = s.chars();
124
125    // First character to lowercase
126    if let Some(first) = chars.next() {
127        result.push(first.to_lowercase().next().unwrap_or(first));
128    }
129
130    // Rest of the string unchanged
131    result.extend(chars);
132    result
133}
134
135/// Convert PascalCase to snake_case
136fn to_snake_case(s: &str) -> String {
137    let mut result = String::new();
138    let chars: Vec<char> = s.chars().collect();
139
140    for (i, &c) in chars.iter().enumerate() {
141        if c.is_uppercase() {
142            // Add underscore before uppercase letter if:
143            // 1. Not the first character
144            // 2. Previous character was lowercase, OR
145            // 3. Next character is lowercase (handles consecutive caps like "XMLHttp" -> "xml_http")
146            if i > 0 {
147                let prev_was_lower = chars[i - 1].is_lowercase();
148                let next_is_lower = i + 1 < chars.len() && chars[i + 1].is_lowercase();
149
150                if prev_was_lower || next_is_lower {
151                    result.push('_');
152                }
153            }
154            result.push(c.to_lowercase().next().unwrap_or(c));
155        } else {
156            result.push(c);
157        }
158    }
159
160    result
161}
162
163/// Convert PascalCase to kebab-case
164fn to_kebab_case(s: &str) -> String {
165    to_snake_case(s).replace('_', "-")
166}
167
168/// Convert PascalCase to SCREAMING_SNAKE_CASE
169fn to_screaming_snake_case(s: &str) -> String {
170    to_snake_case(s).to_uppercase()
171}
172
173/// Convert to PascalCase (already in PascalCase, so return as-is)
174fn to_pascal_case(s: &str) -> String {
175    s.to_string()
176}
177
178/// Convert PascalCase to SCREAMING-KEBAB-CASE
179fn to_screaming_kebab_case(s: &str) -> String {
180    to_kebab_case(s).to_uppercase()
181}
182
183#[cfg(test)]
184mod tests {
185    use super::*;
186
187    #[test]
188    fn test_camel_case() {
189        assert_eq!(to_camel_case("MessageOne"), "messageOne");
190        assert_eq!(to_camel_case("MessageTwo"), "messageTwo");
191        assert_eq!(to_camel_case("A"), "a");
192        assert_eq!(to_camel_case(""), "");
193    }
194
195    #[test]
196    fn test_snake_case() {
197        assert_eq!(to_snake_case("MessageOne"), "message_one");
198        assert_eq!(to_snake_case("MessageTwo"), "message_two");
199        assert_eq!(to_snake_case("XMLHttpRequest"), "xml_http_request");
200    }
201
202    #[test]
203    fn test_kebab_case() {
204        assert_eq!(to_kebab_case("MessageOne"), "message-one");
205        assert_eq!(to_kebab_case("MessageTwo"), "message-two");
206    }
207
208    #[test]
209    fn test_screaming_snake_case() {
210        assert_eq!(to_screaming_snake_case("MessageOne"), "MESSAGE_ONE");
211        assert_eq!(to_screaming_snake_case("MessageTwo"), "MESSAGE_TWO");
212    }
213}