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 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 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 let rename_case = extract_serde_rename_all(&input.attrs); 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
80fn 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 let tokens_str = meta_list.tokens.to_string();
87 if tokens_str.contains("rename_all") {
88 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
103fn 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
116fn 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 if let Some(first) = chars.next() {
127 result.push(first.to_lowercase().next().unwrap_or(first));
128 }
129
130 result.extend(chars);
132 result
133}
134
135fn 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 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
163fn to_kebab_case(s: &str) -> String {
165 to_snake_case(s).replace('_', "-")
166}
167
168fn to_screaming_snake_case(s: &str) -> String {
170 to_snake_case(s).to_uppercase()
171}
172
173fn to_pascal_case(s: &str) -> String {
175 s.to_string()
176}
177
178fn 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}