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    errors::Result,
12    types::{OPTION, VEC, try_extract_wrapper},
13};
14
15// ---------------------------------------------------------------------------
16// Public entry point
17// ---------------------------------------------------------------------------
18
19/// Generate the `#[derive(OrpcError)]` expansion.
20pub 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
92// ---------------------------------------------------------------------------
93// Type → Zod schema string
94//
95// Uses AST-based wrapper detection for Option<T> and Vec<T>.
96// ---------------------------------------------------------------------------
97
98fn zod_schema_for_type(ty: &syn::Type) -> String {
99    // Option<T>
100    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    // Vec<T>
108    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    // Path types — check final segment ident
116    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            // Custom types without #[derive(ZodTs)] cannot be introspected here —
127            // the macro only sees the field's type name, not its internal fields.
128            // Emit z.unknown() so the contract is always valid TypeScript.
129            // To get a precise schema, add #[derive(ZodTs)] to the inner type.
130            _other => "z.unknown()".to_string(),
131        };
132    }
133
134    // Unit type ()
135    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
144// ---------------------------------------------------------------------------
145// Naming helpers
146// ---------------------------------------------------------------------------
147
148/// Convert `PascalCase` to `SCREAMING_SNAKE_CASE`.
149fn 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// ---------------------------------------------------------------------------
163// Tests
164// ---------------------------------------------------------------------------
165
166#[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        // core::option::Option — matches on final segment
206        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        // Custom types without #[derive(ZodTs)] cannot be introspected at macro time —
219        // the macro only sees the type name, not its internal fields.
220        // Falls back to z.unknown() so generated contracts always compile.
221        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}