Skip to main content

rorpc_parse/codegen/
zod_ts.rs

1//! Code generation for `#[derive(ZodTs)]`.
2//!
3//! Generates a `fn zod_ts() -> String` method that returns a complete
4//! TypeScript block with a Zod schema and a `z.infer` type alias,
5//! plus an `inventory::submit!` for `SchemaRegistration`.
6
7use proc_macro2::TokenStream;
8use quote::quote;
9use syn::{Data, DeriveInput, Fields};
10
11use crate::{
12    attributes::{ZodAttrs, apply_rename_rule, parse_serde_attrs, parse_zod_attrs},
13    errors::Result,
14    types::{OPTION, VEC, is_primitive, try_extract_wrapper},
15};
16
17const ZOD_IMPORT: &str = "import * as z from \"zod\";";
18
19// ---------------------------------------------------------------------------
20// Public entry point
21// ---------------------------------------------------------------------------
22
23/// Generate the `#[derive(ZodTs)]` expansion.
24pub fn derive_zod_ts(input: DeriveInput) -> Result<TokenStream> {
25    let name = &input.ident;
26    let name_str = name.to_string();
27
28    match &input.data {
29        Data::Struct(data) => match &data.fields {
30            Fields::Named(fields) => expand_named_struct(name, &name_str, fields, &input),
31            Fields::Unnamed(_) | Fields::Unit => Err(syn::Error::new_spanned(
32                name,
33                "ZodTs: only structs with named fields are supported",
34            )
35            .into()),
36        },
37        Data::Enum(data) => expand_enum(name, &name_str, data, &input),
38        Data::Union(_) => {
39            Err(syn::Error::new_spanned(name, "ZodTs cannot be derived for unions").into())
40        }
41    }
42}
43
44// ---------------------------------------------------------------------------
45// Struct expansion
46// ---------------------------------------------------------------------------
47
48fn expand_named_struct(
49    name: &syn::Ident,
50    name_str: &str,
51    fields: &syn::FieldsNamed,
52    _input: &DeriveInput,
53) -> Result<TokenStream> {
54    let mut field_lines: Vec<String> = Vec::new();
55    let mut dep_type_names: Vec<String> = Vec::new();
56
57    for field in &fields.named {
58        let field_name = field.ident.as_ref().unwrap().to_string();
59        let serde = parse_serde_attrs(&field.attrs)?;
60
61        if serde.skip {
62            continue;
63        }
64
65        let ts_key = serde.rename.as_deref().unwrap_or(&field_name);
66        let zod = parse_zod_attrs(&field.attrs)?;
67        let is_opt = is_option_type(&field.ty);
68
69        let base_ty = if is_opt {
70            option_inner(&field.ty).unwrap_or(&field.ty)
71        } else {
72            &field.ty
73        };
74
75        let zod_expr = rust_type_to_zod(base_ty, &zod);
76
77        let final_expr = if is_opt {
78            format!("{}.optional()", zod_expr)
79        } else {
80            zod_expr
81        };
82
83        field_lines.push(format!("  {}: {}", ts_key, final_expr));
84
85        // Collect non-primitive custom types for dependency tracking
86        if let Some(custom) = innermost_custom_name(base_ty) {
87            dep_type_names.push(custom);
88        }
89    }
90
91    let schema_name = format!("{}Schema", name_str);
92    let ts_code = format!(
93        "{}\n\nexport const {} = z.object({{\n{}\n}});\n\nexport type {} = z.infer<typeof {}>;",
94        ZOD_IMPORT,
95        schema_name,
96        field_lines.join(",\n"),
97        name_str,
98        schema_name,
99    );
100
101    Ok(emit_registration(name, name_str, &ts_code, &dep_type_names))
102}
103
104// ---------------------------------------------------------------------------
105// Enum expansion
106// ---------------------------------------------------------------------------
107
108fn expand_enum(
109    name: &syn::Ident,
110    name_str: &str,
111    data: &syn::DataEnum,
112    input: &DeriveInput,
113) -> Result<TokenStream> {
114    let serde_container = parse_serde_attrs(&input.attrs)?;
115    let rename_all = serde_container.rename_all.as_deref();
116
117    let mut variant_schemas: Vec<String> = Vec::new();
118
119    for variant in &data.variants {
120        let serde_variant = parse_serde_attrs(&variant.attrs)?;
121        if serde_variant.skip {
122            continue;
123        }
124
125        let raw_name = variant.ident.to_string();
126        let variant_name = serde_variant
127            .rename
128            .as_deref()
129            .map(str::to_string)
130            .unwrap_or_else(|| {
131                rename_all
132                    .map(|rule| apply_rename_rule(rule, &raw_name))
133                    .unwrap_or(raw_name)
134            });
135
136        variant_schemas.push(generate_variant_ts(&variant_name, &variant.fields)?);
137    }
138
139    let schema_name = format!("{}Schema", name_str);
140    let variants_str = variant_schemas.join(",\n  ");
141    let ts_code = format!(
142        "{}\n\nexport const {} = z.union([\n  {}\n]);\n\nexport type {} = z.infer<typeof {}>;",
143        ZOD_IMPORT, schema_name, variants_str, name_str, schema_name,
144    );
145
146    Ok(emit_registration(name, name_str, &ts_code, &[]))
147}
148
149// ---------------------------------------------------------------------------
150// Variant code generation
151// ---------------------------------------------------------------------------
152
153fn generate_variant_ts(variant_name: &str, fields: &Fields) -> Result<String> {
154    match fields {
155        Fields::Unit => Ok(format!("z.literal(\"{}\")", escape_str(variant_name))),
156
157        Fields::Unnamed(fields_unnamed) => {
158            let count = fields_unnamed.unnamed.len();
159            if count == 1 {
160                let field = fields_unnamed.unnamed.first().unwrap();
161                let zod = parse_zod_attrs(&field.attrs)?;
162                let schema = rust_type_to_zod(&field.ty, &zod);
163                Ok(format!(
164                    "z.object({{ {}: {} }})",
165                    ts_object_key(variant_name),
166                    schema
167                ))
168            } else {
169                let elements: Vec<String> = fields_unnamed
170                    .unnamed
171                    .iter()
172                    .map(|f| {
173                        let zod = parse_zod_attrs(&f.attrs)?;
174                        Ok(rust_type_to_zod(&f.ty, &zod))
175                    })
176                    .collect::<Result<Vec<_>>>()?;
177                Ok(format!(
178                    "z.object({{ {}: z.tuple([{}]) }})",
179                    ts_object_key(variant_name),
180                    elements.join(", ")
181                ))
182            }
183        }
184
185        Fields::Named(fields_named) => {
186            let field_schemas: Vec<String> = fields_named
187                .named
188                .iter()
189                .map(|field| {
190                    let field_name = field.ident.as_ref().unwrap().to_string();
191                    let serde = parse_serde_attrs(&field.attrs)?;
192                    if serde.skip {
193                        return Ok(String::new());
194                    }
195                    let ts_key = serde.rename.as_deref().unwrap_or(&field_name);
196                    let zod_attrs = parse_zod_attrs(&field.attrs)?;
197                    let is_opt = is_option_type(&field.ty);
198                    let base_ty = if is_opt {
199                        option_inner(&field.ty).unwrap_or(&field.ty)
200                    } else {
201                        &field.ty
202                    };
203                    let schema = rust_type_to_zod(base_ty, &zod_attrs);
204                    let final_schema = if is_opt {
205                        format!("{}.optional()", schema)
206                    } else {
207                        schema
208                    };
209                    Ok(format!("{}: {}", ts_key, final_schema))
210                })
211                .filter(|r| r.as_deref().map(|s| !s.is_empty()).unwrap_or(true))
212                .collect::<Result<Vec<_>>>()?;
213
214            Ok(format!(
215                "z.object({{ {}: z.object({{ {} }}) }})",
216                ts_object_key(variant_name),
217                field_schemas.join(", ")
218            ))
219        }
220    }
221}
222
223// ---------------------------------------------------------------------------
224// inventory::submit! emission
225// ---------------------------------------------------------------------------
226
227fn emit_registration(
228    name: &syn::Ident,
229    name_str: &str,
230    ts_code: &str,
231    dep_type_names: &[String],
232) -> TokenStream {
233    let dep_strs: Vec<&str> = dep_type_names.iter().map(String::as_str).collect();
234
235    quote! {
236        impl #name {
237            pub fn zod_ts() -> String {
238                #ts_code.to_string()
239            }
240
241            pub fn dependent_types() -> Vec<&'static str> {
242                vec![#(#dep_strs),*]
243            }
244        }
245
246        const _: () = {
247            ::rorpc::inventory::submit! {
248                ::rorpc::SchemaRegistration {
249                    type_name: #name_str,
250                    zod_ts: #name::zod_ts,
251                    dependent_types: #name::dependent_types,
252                }
253            }
254        };
255    }
256}
257
258// ---------------------------------------------------------------------------
259// Type → Zod expression
260// ---------------------------------------------------------------------------
261
262/// Map a `syn::Type` to a Zod schema expression string.
263///
264/// Uses AST-based wrapper detection for `Option<T>` and `Vec<T>` —
265/// never string prefix matching.
266pub fn rust_type_to_zod(ty: &syn::Type, attrs: &ZodAttrs) -> String {
267    // Option<T> — recurse on inner, then .optional()
268    if is_option_type(ty)
269        && let Some(inner) = option_inner(ty)
270    {
271        let inner_schema = rust_type_to_zod(inner, &ZodAttrs::default());
272        return format!("{}.optional()", inner_schema);
273    }
274
275    // Vec<T>
276    if let Some(m) = try_extract_wrapper(ty, VEC)
277        && let Some(inner) = m.first_type()
278    {
279        let inner_schema = rust_type_to_zod(inner, &ZodAttrs::default());
280        let mut chain = format!("z.array({})", inner_schema);
281        if let Some(n) = attrs.length {
282            chain.push_str(&format!(".length({})", n));
283        }
284        if let Some(n) = attrs.min_length {
285            chain.push_str(&format!(".min({})", n));
286        }
287        if let Some(n) = attrs.max_length {
288            chain.push_str(&format!(".max({})", n));
289        }
290        return chain;
291    }
292
293    // Primitives — match on the final path segment ident
294    if let syn::Type::Path(type_path) = ty
295        && let Some(seg) = type_path.path.segments.last()
296    {
297        let name = seg.ident.to_string();
298        return match name.as_str() {
299            "String" | "str" => build_string_schema(attrs),
300            "i8" | "i16" | "i32" | "i64" | "i128" | "isize" | "u8" | "u16" | "u32" | "u64"
301            | "u128" | "usize" => build_integer_schema(attrs),
302            "f32" | "f64" => build_float_schema(attrs),
303            "bool" => "z.boolean()".to_string(),
304            // serde_json::Value → z.any()
305            "Value" => "z.any()".to_string(),
306            // Custom type — reference its schema by name
307            other => format!("{}Schema", other),
308        };
309    }
310
311    // Unit type ()
312    if let syn::Type::Tuple(t) = ty
313        && t.elems.is_empty()
314    {
315        return "z.void()".to_string();
316    }
317
318    "z.unknown()".to_string()
319}
320
321// ---------------------------------------------------------------------------
322// Schema builders
323// ---------------------------------------------------------------------------
324
325fn build_string_schema(attrs: &ZodAttrs) -> String {
326    let mut chain = String::from("z.string()");
327    if let Some(n) = attrs.length {
328        chain.push_str(&format!(".length({})", n));
329    }
330    if let Some(n) = attrs.min_length {
331        chain.push_str(&format!(".min({})", n));
332    }
333    if let Some(n) = attrs.max_length {
334        chain.push_str(&format!(".max({})", n));
335    }
336    if attrs.email {
337        chain.push_str(".email()");
338    }
339    if attrs.url {
340        chain.push_str(".url()");
341    }
342    if let Some(ref p) = attrs.regex {
343        chain.push_str(&format!(".regex(/{}/)", p));
344    }
345    if let Some(ref p) = attrs.starts_with {
346        chain.push_str(&format!(".startsWith(\"{}\")", p));
347    }
348    if let Some(ref p) = attrs.ends_with {
349        chain.push_str(&format!(".endsWith(\"{}\")", p));
350    }
351    if let Some(ref p) = attrs.includes {
352        chain.push_str(&format!(".includes(\"{}\")", p));
353    }
354    chain
355}
356
357fn build_integer_schema(attrs: &ZodAttrs) -> String {
358    let mut chain = String::from("z.number().int()");
359    append_number_validators(&mut chain, attrs);
360    chain
361}
362
363fn build_float_schema(attrs: &ZodAttrs) -> String {
364    let mut chain = String::from("z.number()");
365    if attrs.int {
366        chain.push_str(".int()");
367    }
368    append_number_validators(&mut chain, attrs);
369    chain
370}
371
372fn append_number_validators(chain: &mut String, attrs: &ZodAttrs) {
373    if let Some(n) = attrs.min {
374        chain.push_str(&format!(".min({})", n));
375    }
376    if let Some(n) = attrs.max {
377        chain.push_str(&format!(".max({})", n));
378    }
379    if attrs.positive {
380        chain.push_str(".positive()");
381    }
382    if attrs.negative {
383        chain.push_str(".negative()");
384    }
385    if attrs.nonnegative {
386        chain.push_str(".nonnegative()");
387    }
388    if attrs.nonpositive {
389        chain.push_str(".nonpositive()");
390    }
391    if attrs.finite {
392        chain.push_str(".finite()");
393    }
394}
395
396// ---------------------------------------------------------------------------
397// Type helpers — all AST-based, no string matching on type names
398// ---------------------------------------------------------------------------
399
400fn is_option_type(ty: &syn::Type) -> bool {
401    try_extract_wrapper(ty, OPTION).is_some()
402}
403
404fn option_inner(ty: &syn::Type) -> Option<&syn::Type> {
405    try_extract_wrapper(ty, OPTION)?.first_type()
406}
407
408/// Return the simple name of the innermost non-primitive, non-wrapper type,
409/// for dependency tracking in `dependent_types()`.
410fn innermost_custom_name(ty: &syn::Type) -> Option<String> {
411    // Strip Vec<T>
412    if let Some(m) = try_extract_wrapper(ty, VEC) {
413        return m.first_type().and_then(innermost_custom_name);
414    }
415    if is_primitive(ty) {
416        return None;
417    }
418    if let syn::Type::Path(tp) = ty
419        && let Some(seg) = tp.path.segments.last()
420    {
421        let name = seg.ident.to_string();
422        // Exclude Value (serde_json) from dependency tracking
423        if name == "Value" {
424            return None;
425        }
426        return Some(name);
427    }
428    None
429}
430
431fn ts_object_key(name: &str) -> String {
432    let valid = !name.is_empty()
433        && name
434            .chars()
435            .next()
436            .is_some_and(|c| c.is_ascii_alphabetic() || c == '_' || c == '$')
437        && name
438            .chars()
439            .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '$');
440    if valid {
441        name.to_string()
442    } else {
443        format!("\"{}\"", escape_str(name))
444    }
445}
446
447fn escape_str(s: &str) -> String {
448    s.replace('\\', "\\\\").replace('"', "\\\"")
449}
450
451// ---------------------------------------------------------------------------
452// Runtime type-to-zod conversion (for contract generation)
453// ---------------------------------------------------------------------------
454
455/// Convert a Rust type name string to a TypeScript Zod schema reference.
456///
457/// This is for runtime contract generation when you have type names as strings
458/// from handler metadata, not `syn::Type` ASTs. For compile-time AST-based
459/// conversion, use [`rust_type_to_zod`] instead.
460///
461/// # String-based parsing
462///
463/// This function uses string prefix/suffix matching because it operates on
464/// type name strings collected at link time via `inventory`. It handles:
465///
466/// - Wrapper unwrapping: `"Json<Planet>"` → `"PlanetSchema"`
467/// - Result unwrapping: `"Result<Json<Planet>, E>"` → `"PlanetSchema"`
468/// - Vec mapping: `"Json<Vec<Planet>>"` → `"z.array(PlanetSchema)"`
469/// - Primitive mapping: `"String"` → `"z.string()"`
470/// - SSE streams: `"Sse<...>"` → `"asyncIteratorObject(z.unknown())"`
471///
472/// # Examples
473///
474/// ```
475/// use rorpc_parse::codegen::rust_type_to_ts_schema;
476///
477/// assert_eq!(rust_type_to_ts_schema("Json<Planet>"), "PlanetSchema");
478/// assert_eq!(rust_type_to_ts_schema("Json<Vec<Planet>>"), "z.array(PlanetSchema)");
479/// assert_eq!(rust_type_to_ts_schema("Result<Json<Planet>, E>"), "PlanetSchema");
480/// assert_eq!(rust_type_to_ts_schema("String"), "z.string()");
481/// assert_eq!(rust_type_to_ts_schema("()"), "");
482/// ```
483pub fn rust_type_to_ts_schema(raw: &str) -> String {
484    let raw = raw.replace(' ', "");
485
486    if raw.starts_with("Sse<") {
487        return "asyncIteratorObject(z.unknown() /* TODO: add #[derive(ZodTs)] to your stream event type */)".to_string();
488    }
489
490    // Unwrap Result<T, E> → T
491    let inner = if raw.starts_with("Result<") {
492        extract_first_generic_arg_string(&raw).unwrap_or(raw.clone())
493    } else {
494        raw.clone()
495    };
496
497    // Unwrap Json<T> → T
498    let inner = if inner.starts_with("Json<") && inner.ends_with('>') {
499        inner[5..inner.len() - 1].to_string()
500    } else {
501        inner
502    };
503
504    type_name_to_zod_ref(&inner)
505}
506
507/// Map a bare type name to its Zod schema reference.
508fn type_name_to_zod_ref(type_name: &str) -> String {
509    match type_name {
510        "()" | "" => String::new(),
511        "String" | "str" => "z.string()".to_string(),
512        "bool" => "z.boolean()".to_string(),
513        "i8" | "i16" | "i32" | "i64" | "i128" | "isize" | "u8" | "u16" | "u32" | "u64" | "u128"
514        | "usize" => "z.number().int()".to_string(),
515        "f32" | "f64" => "z.number()".to_string(),
516        "serde_json::Value" | "Value" => "z.any()".to_string(),
517        _ if type_name.starts_with("Vec<") && type_name.ends_with('>') => {
518            let inner = &type_name[4..type_name.len() - 1];
519            format!("z.array({})", type_name_to_zod_ref(inner))
520        }
521        _ if type_name.starts_with("Option<") && type_name.ends_with('>') => {
522            let inner = &type_name[7..type_name.len() - 1];
523            format!("{}.optional()", type_name_to_zod_ref(inner))
524        }
525        _ => {
526            let base = type_name.rsplit("::").next().unwrap_or(type_name);
527            format!("{}Schema", base)
528        }
529    }
530}
531
532/// Extract the first generic argument from a type string.
533///
534/// `"Result<Json<Planet>, E>"` → `Some("Json<Planet>")`
535fn extract_first_generic_arg_string(type_str: &str) -> Option<String> {
536    let start = type_str.find('<')? + 1;
537    let mut depth = 0;
538    let mut end = start;
539
540    for (i, ch) in type_str[start..].char_indices() {
541        match ch {
542            '<' => depth += 1,
543            '>' if depth == 0 => {
544                end = start + i;
545                break;
546            }
547            '>' => depth -= 1,
548            ',' if depth == 0 => {
549                end = start + i;
550                break;
551            }
552            _ => {}
553        }
554    }
555
556    if end > start {
557        Some(type_str[start..end].to_string())
558    } else {
559        None
560    }
561}
562
563/// Convert type name to schema constant name: `"Planet"` → `"PlanetSchema"`
564pub fn to_schema_name(rust_type: &str) -> String {
565    format!("{}Schema", base_type_name(rust_type))
566}
567
568/// Extract the base type name, stripping all wrappers.
569///
570/// `"Result<Json<Vec<Planet>>, E>"` → `"Planet"`
571pub fn base_type_name(rust_type: &str) -> String {
572    let mut base = rust_type.trim();
573
574    if base.starts_with("Result<")
575        && let Some(inner) = extract_first_generic_arg_string(base)
576    {
577        base = Box::leak(inner.into_boxed_str());
578    }
579    if base.starts_with("Json<") && base.ends_with('>') {
580        base = &base[5..base.len() - 1];
581    }
582    if base.starts_with("Vec<") && base.ends_with('>') {
583        base = &base[4..base.len() - 1];
584    }
585    if base.starts_with("Option<") && base.ends_with('>') {
586        base = &base[7..base.len() - 1];
587    }
588
589    base.rsplit("::").next().unwrap_or(base).to_string()
590}
591
592#[cfg(test)]
593mod runtime_conversion_tests {
594    use super::*;
595
596    #[test]
597    fn json_planet() {
598        assert_eq!(rust_type_to_ts_schema("Json<Planet>"), "PlanetSchema");
599    }
600
601    #[test]
602    fn json_vec_planet() {
603        assert_eq!(
604            rust_type_to_ts_schema("Json<Vec<Planet>>"),
605            "z.array(PlanetSchema)"
606        );
607    }
608
609    #[test]
610    fn result_json_planet() {
611        assert_eq!(
612            rust_type_to_ts_schema("Result<Json<Planet>, StatusCode>"),
613            "PlanetSchema"
614        );
615    }
616
617    #[test]
618    fn json_string() {
619        assert_eq!(rust_type_to_ts_schema("Json<String>"), "z.string()");
620    }
621
622    #[test]
623    fn unit_type() {
624        assert_eq!(rust_type_to_ts_schema("()"), "");
625    }
626
627    #[test]
628    fn serde_json_value() {
629        assert_eq!(rust_type_to_ts_schema("Json<serde_json::Value>"), "z.any()");
630    }
631
632    #[test]
633    fn schema_name_simple() {
634        assert_eq!(to_schema_name("Planet"), "PlanetSchema");
635    }
636
637    #[test]
638    fn schema_name_vec() {
639        assert_eq!(to_schema_name("Vec<Planet>"), "PlanetSchema");
640    }
641
642    #[test]
643    fn base_type_unwraps_wrappers() {
644        assert_eq!(base_type_name("Result<Json<Vec<Planet>>, E>"), "Planet");
645        assert_eq!(base_type_name("Json<Planet>"), "Planet");
646        assert_eq!(base_type_name("Vec<Planet>"), "Planet");
647        assert_eq!(base_type_name("Option<Planet>"), "Planet");
648    }
649
650    #[test]
651    fn base_type_strips_module_path() {
652        assert_eq!(base_type_name("models::Planet"), "Planet");
653        assert_eq!(base_type_name("crate::domain::Planet"), "Planet");
654    }
655}