rorpc_parse/codegen/
error_derive.rs1use proc_macro2::TokenStream;
7use quote::quote;
8use syn::{Data, DeriveInput, Fields};
9
10use crate::{
11 errors::Result,
12 types::{OPTION, VEC, try_extract_wrapper},
13};
14
15pub fn expand_orpc_errors(input: DeriveInput) -> Result<TokenStream> {
21 let enum_name = &input.ident;
22 let enum_name_str = enum_name.to_string();
23
24 let variants = match &input.data {
25 Data::Enum(data) => &data.variants,
26 _ => {
27 return Err(syn::Error::new_spanned(
28 &input.ident,
29 "#[derive(OrpcError)] can only be used on enums",
30 )
31 .into());
32 }
33 };
34
35 let mut variant_tokens: Vec<TokenStream> = Vec::new();
36
37 for variant in variants {
38 let variant_name_screaming = to_screaming_snake_case(&variant.ident.to_string());
39
40 let data_schema = match &variant.fields {
41 Fields::Unit => quote! { None },
42 Fields::Unnamed(fields) => {
43 let schemas: Vec<String> = fields
44 .unnamed
45 .iter()
46 .map(|f| zod_schema_for_type(&f.ty))
47 .collect();
48 let schema = if schemas.len() == 1 {
49 schemas.into_iter().next().unwrap()
50 } else {
51 format!("z.tuple([{}])", schemas.join(", "))
52 };
53 quote! { Some(#schema) }
54 }
55 Fields::Named(fields) => {
56 let field_schemas: Vec<String> = fields
57 .named
58 .iter()
59 .map(|f| {
60 let name = f.ident.as_ref().unwrap().to_string();
61 let schema = zod_schema_for_type(&f.ty);
62 format!("{}: {}", name, schema)
63 })
64 .collect();
65 let schema = format!("z.object({{ {} }})", field_schemas.join(", "));
66 quote! { Some(#schema) }
67 }
68 };
69
70 variant_tokens.push(quote! {
71 ::rorpc::ErrorVariant {
72 name: #variant_name_screaming,
73 data_schema: #data_schema,
74 }
75 });
76 }
77
78 Ok(quote! {
79 const _: () = {
80 ::rorpc::inventory::submit! {
81 ::rorpc::ErrorRegistration {
82 type_name: #enum_name_str,
83 variants: &[
84 #(#variant_tokens),*
85 ],
86 }
87 }
88 };
89 })
90}
91
92fn zod_schema_for_type(ty: &syn::Type) -> String {
99 if let Some(m) = try_extract_wrapper(ty, OPTION) {
101 if let Some(inner) = m.first_type() {
102 return format!("{}.optional()", zod_schema_for_type(inner));
103 }
104 return "z.unknown().optional()".to_string();
105 }
106
107 if let Some(m) = try_extract_wrapper(ty, VEC) {
109 if let Some(inner) = m.first_type() {
110 return format!("z.array({})", zod_schema_for_type(inner));
111 }
112 return "z.array(z.unknown())".to_string();
113 }
114
115 if let syn::Type::Path(type_path) = ty
117 && let Some(seg) = type_path.path.segments.last()
118 {
119 return match seg.ident.to_string().as_str() {
120 "String" | "str" => "z.string()".to_string(),
121 "i8" | "i16" | "i32" | "i64" | "i128" | "isize" | "u8" | "u16" | "u32" | "u64"
122 | "u128" | "usize" => "z.number().int()".to_string(),
123 "f32" | "f64" => "z.number()".to_string(),
124 "bool" => "z.boolean()".to_string(),
125 "Value" => "z.record(z.string(), z.unknown())".to_string(),
126 _other => "z.unknown()".to_string(),
131 };
132 }
133
134 if let syn::Type::Tuple(t) = ty
136 && t.elems.is_empty()
137 {
138 return "z.void()".to_string();
139 }
140
141 "z.unknown()".to_string()
142}
143
144fn to_screaming_snake_case(name: &str) -> String {
150 let mut out = String::new();
151 let mut prev_lower = false;
152 for (i, ch) in name.chars().enumerate() {
153 if ch.is_uppercase() && i > 0 && prev_lower {
154 out.push('_');
155 }
156 out.push(ch.to_ascii_uppercase());
157 prev_lower = ch.is_ascii_lowercase();
158 }
159 out
160}
161
162#[cfg(test)]
167mod tests {
168 use super::*;
169
170 #[test]
171 fn screaming_snake_case() {
172 assert_eq!(to_screaming_snake_case("NotFound"), "NOT_FOUND");
173 assert_eq!(to_screaming_snake_case("DatabaseError"), "DATABASE_ERROR");
174 assert_eq!(
175 to_screaming_snake_case("RateLimitExceeded"),
176 "RATE_LIMIT_EXCEEDED"
177 );
178 assert_eq!(to_screaming_snake_case("OK"), "OK");
179 assert_eq!(to_screaming_snake_case("NotFoundError"), "NOT_FOUND_ERROR");
180 }
181
182 #[test]
183 fn zod_schema_primitives() {
184 let string_ty: syn::Type = syn::parse_str("String").unwrap();
185 assert_eq!(zod_schema_for_type(&string_ty), "z.string()");
186
187 let i32_ty: syn::Type = syn::parse_str("i32").unwrap();
188 assert_eq!(zod_schema_for_type(&i32_ty), "z.number().int()");
189
190 let f64_ty: syn::Type = syn::parse_str("f64").unwrap();
191 assert_eq!(zod_schema_for_type(&f64_ty), "z.number()");
192
193 let bool_ty: syn::Type = syn::parse_str("bool").unwrap();
194 assert_eq!(zod_schema_for_type(&bool_ty), "z.boolean()");
195 }
196
197 #[test]
198 fn zod_schema_option() {
199 let ty: syn::Type = syn::parse_str("Option<String>").unwrap();
200 assert_eq!(zod_schema_for_type(&ty), "z.string().optional()");
201 }
202
203 #[test]
204 fn zod_schema_qualified_option() {
205 let ty: syn::Type = syn::parse_str("core::option::Option<i32>").unwrap();
207 assert_eq!(zod_schema_for_type(&ty), "z.number().int().optional()");
208 }
209
210 #[test]
211 fn zod_schema_vec() {
212 let ty: syn::Type = syn::parse_str("Vec<String>").unwrap();
213 assert_eq!(zod_schema_for_type(&ty), "z.array(z.string())");
214 }
215
216 #[test]
217 fn zod_schema_custom_type() {
218 let ty: syn::Type = syn::parse_str("Planet").unwrap();
222 assert_eq!(zod_schema_for_type(&ty), "z.unknown()");
223 }
224
225 #[test]
226 fn zod_schema_value() {
227 let ty: syn::Type = syn::parse_str("serde_json::Value").unwrap();
228 assert_eq!(
229 zod_schema_for_type(&ty),
230 "z.record(z.string(), z.unknown())"
231 );
232 }
233
234 #[test]
235 fn expand_unit_enum() {
236 let input: DeriveInput = syn::parse_quote! {
237 enum AppError { NotFound, Conflict }
238 };
239 let ts = expand_orpc_errors(input).unwrap();
240 let code = ts.to_string();
241 assert!(code.contains("NOT_FOUND"));
242 assert!(code.contains("CONFLICT"));
243 assert!(code.contains("None"));
244 }
245
246 #[test]
247 fn expand_rejects_struct() {
248 let input: DeriveInput = syn::parse_quote! {
249 struct NotAnEnum { field: String }
250 };
251 assert!(expand_orpc_errors(input).is_err());
252 }
253}