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