Skip to main content

rorpc_parse/codegen/
error_derive.rs

1//! Code generation for `#[derive(OrpcError)]`.
2//!
3//! Registers error enum variants with rorpc via `inventory::submit!` so
4//! `generate_contract()` can emit TypeScript `.errors({...})` entries.
5
6use 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
16// ---------------------------------------------------------------------------
17// Public entry point
18// ---------------------------------------------------------------------------
19
20/// Generate the `#[derive(OrpcError)]` expansion.
21pub 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
93// ---------------------------------------------------------------------------
94// Type → Zod schema string
95//
96// Uses AST-based wrapper detection for Option<T> and Vec<T>.
97// ---------------------------------------------------------------------------
98
99fn zod_schema_for_type(ty: &syn::Type) -> String {
100    // Option<T>
101    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    // Vec<T>
109    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    // HashMap<K, V>
117    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    // Path types — check final segment ident
128    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        // Check primitives using shared helper
134        if let Some(zod) = primitive_zod_expr(&type_name) {
135            return zod.to_string();
136        }
137
138        // Custom types without #[derive(ZodTs)] cannot be introspected here —
139        // the macro only sees the field's type name, not its internal fields.
140        // Emit z.unknown() so the contract is always valid TypeScript.
141        // To get a precise schema, add #[derive(ZodTs)] to the inner type.
142        return "z.unknown()".to_string();
143    }
144
145    // Unit type ()
146    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
155// ---------------------------------------------------------------------------
156// Naming helpers
157// ---------------------------------------------------------------------------
158
159/// Convert `PascalCase` to `SCREAMING_SNAKE_CASE`.
160fn 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// ---------------------------------------------------------------------------
174// Tests
175// ---------------------------------------------------------------------------
176
177#[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        // core::option::Option — matches on final segment
217        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        // Custom types without #[derive(ZodTs)] cannot be introspected at macro time —
245        // the macro only sees the type name, not its internal fields.
246        // Falls back to z.unknown() so generated contracts always compile.
247        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}