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            // uuid::Uuid → z.uuid()
305            "Uuid" => "z.uuid()".to_string(),
306            // chrono::DateTime<Utc> → z.iso.datetime()
307            "DateTime" => "z.iso.datetime({ offset: true })".to_string(),
308            // serde_json::Value → z.any()
309            "Value" => "z.any()".to_string(),
310            // Custom type — reference its schema by name
311            other => format!("{}Schema", other),
312        };
313    }
314
315    // Unit type ()
316    if let syn::Type::Tuple(t) = ty
317        && t.elems.is_empty()
318    {
319        return "z.void()".to_string();
320    }
321
322    "z.unknown()".to_string()
323}
324
325// ---------------------------------------------------------------------------
326// Schema builders
327// ---------------------------------------------------------------------------
328
329fn build_string_schema(attrs: &ZodAttrs) -> String {
330    let mut chain = String::from("z.string()");
331    if let Some(n) = attrs.length {
332        chain.push_str(&format!(".length({})", n));
333    }
334    if let Some(n) = attrs.min_length {
335        chain.push_str(&format!(".min({})", n));
336    }
337    if let Some(n) = attrs.max_length {
338        chain.push_str(&format!(".max({})", n));
339    }
340    if attrs.email {
341        chain.push_str(".email()");
342    }
343    if attrs.url {
344        chain.push_str(".url()");
345    }
346    if let Some(ref p) = attrs.regex {
347        chain.push_str(&format!(".regex(/{}/)", p));
348    }
349    if let Some(ref p) = attrs.starts_with {
350        chain.push_str(&format!(".startsWith(\"{}\")", p));
351    }
352    if let Some(ref p) = attrs.ends_with {
353        chain.push_str(&format!(".endsWith(\"{}\")", p));
354    }
355    if let Some(ref p) = attrs.includes {
356        chain.push_str(&format!(".includes(\"{}\")", p));
357    }
358    chain
359}
360
361fn build_integer_schema(attrs: &ZodAttrs) -> String {
362    let mut chain = String::from("z.number().int()");
363    append_number_validators(&mut chain, attrs);
364    chain
365}
366
367fn build_float_schema(attrs: &ZodAttrs) -> String {
368    let mut chain = String::from("z.number()");
369    if attrs.int {
370        chain.push_str(".int()");
371    }
372    append_number_validators(&mut chain, attrs);
373    chain
374}
375
376fn append_number_validators(chain: &mut String, attrs: &ZodAttrs) {
377    if let Some(n) = attrs.min {
378        chain.push_str(&format!(".min({})", n));
379    }
380    if let Some(n) = attrs.max {
381        chain.push_str(&format!(".max({})", n));
382    }
383    if attrs.positive {
384        chain.push_str(".positive()");
385    }
386    if attrs.negative {
387        chain.push_str(".negative()");
388    }
389    if attrs.nonnegative {
390        chain.push_str(".nonnegative()");
391    }
392    if attrs.nonpositive {
393        chain.push_str(".nonpositive()");
394    }
395    if attrs.finite {
396        chain.push_str(".finite()");
397    }
398}
399
400// ---------------------------------------------------------------------------
401// Type helpers — all AST-based, no string matching on type names
402// ---------------------------------------------------------------------------
403
404fn is_option_type(ty: &syn::Type) -> bool {
405    try_extract_wrapper(ty, OPTION).is_some()
406}
407
408fn option_inner(ty: &syn::Type) -> Option<&syn::Type> {
409    try_extract_wrapper(ty, OPTION)?.first_type()
410}
411
412/// Return the simple name of the innermost non-primitive, non-wrapper type,
413/// for dependency tracking in `dependent_types()`.
414fn innermost_custom_name(ty: &syn::Type) -> Option<String> {
415    // Strip Vec<T>
416    if let Some(m) = try_extract_wrapper(ty, VEC) {
417        return m.first_type().and_then(innermost_custom_name);
418    }
419    if is_primitive(ty) {
420        return None;
421    }
422    if let syn::Type::Path(tp) = ty
423        && let Some(seg) = tp.path.segments.last()
424    {
425        let name = seg.ident.to_string();
426        // Exclude Value (serde_json) from dependency tracking
427        if name == "Value" {
428            return None;
429        }
430        return Some(name);
431    }
432    None
433}
434
435fn ts_object_key(name: &str) -> String {
436    let valid = !name.is_empty()
437        && name
438            .chars()
439            .next()
440            .is_some_and(|c| c.is_ascii_alphabetic() || c == '_' || c == '$')
441        && name
442            .chars()
443            .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '$');
444    if valid {
445        name.to_string()
446    } else {
447        format!("\"{}\"", escape_str(name))
448    }
449}
450
451fn escape_str(s: &str) -> String {
452    s.replace('\\', "\\\\").replace('"', "\\\"")
453}
454
455// ---------------------------------------------------------------------------
456// Runtime type-to-zod conversion (for contract generation)
457// ---------------------------------------------------------------------------
458
459/// Convert a Rust type name string to a TypeScript Zod schema reference.
460///
461/// This is for runtime contract generation when you have type names as strings
462/// from handler metadata, not `syn::Type` ASTs. For compile-time AST-based
463/// conversion, use [`rust_type_to_zod`] instead.
464///
465/// # String-based parsing
466///
467/// This function uses string prefix/suffix matching because it operates on
468/// type name strings collected at link time via `inventory`. It handles:
469///
470/// - Wrapper unwrapping: `"Json<Planet>"` → `"PlanetSchema"`
471/// - Result unwrapping: `"Result<Json<Planet>, E>"` → `"PlanetSchema"`
472/// - Vec mapping: `"Json<Vec<Planet>>"` → `"z.array(PlanetSchema)"`
473/// - Primitive mapping: `"String"` → `"z.string()"`
474/// - SSE streams: `"Sse<...>"` → `"asyncIteratorObject(z.unknown())"`
475///
476/// # Examples
477///
478/// ```
479/// use rorpc_parse::codegen::rust_type_to_ts_schema;
480///
481/// assert_eq!(rust_type_to_ts_schema("Json<Planet>"), "PlanetSchema");
482/// assert_eq!(rust_type_to_ts_schema("Json<Vec<Planet>>"), "z.array(PlanetSchema)");
483/// assert_eq!(rust_type_to_ts_schema("Result<Json<Planet>, E>"), "PlanetSchema");
484/// assert_eq!(rust_type_to_ts_schema("String"), "z.string()");
485/// assert_eq!(rust_type_to_ts_schema("()"), "");
486/// ```
487pub fn rust_type_to_ts_schema(raw: &str) -> String {
488    let raw = raw.replace(' ', "");
489
490    if raw.starts_with("Sse<") {
491        return "asyncIteratorObject(z.unknown() /* TODO: add #[derive(ZodTs)] to your stream event type */)".to_string();
492    }
493
494    // Unwrap Result<T, E> → T
495    let inner = if raw.starts_with("Result<") {
496        extract_first_generic_arg_string(&raw).unwrap_or(raw.clone())
497    } else {
498        raw.clone()
499    };
500
501    // Unwrap Json<T> → T
502    let inner = if inner.starts_with("Json<") && inner.ends_with('>') {
503        inner[5..inner.len() - 1].to_string()
504    } else {
505        inner
506    };
507
508    type_name_to_zod_ref(&inner)
509}
510
511/// Map a bare type name to its Zod schema reference.
512fn type_name_to_zod_ref(type_name: &str) -> String {
513    match type_name {
514        "()" | "" => String::new(),
515        "String" | "str" => "z.string()".to_string(),
516        "bool" => "z.boolean()".to_string(),
517        "i8" | "i16" | "i32" | "i64" | "i128" | "isize" | "u8" | "u16" | "u32" | "u64" | "u128"
518        | "usize" => "z.number().int()".to_string(),
519        "f32" | "f64" => "z.number()".to_string(),
520        "Uuid" => "z.uuid()".to_string(),
521        "DateTime" => "z.iso.datetime({ offset: true })".to_string(),
522        "serde_json::Value" | "Value" => "z.any()".to_string(),
523        _ if type_name.starts_with("Vec<") && type_name.ends_with('>') => {
524            let inner = &type_name[4..type_name.len() - 1];
525            format!("z.array({})", type_name_to_zod_ref(inner))
526        }
527        _ if type_name.starts_with("Option<") && type_name.ends_with('>') => {
528            let inner = &type_name[7..type_name.len() - 1];
529            format!("{}.optional()", type_name_to_zod_ref(inner))
530        }
531        _ => {
532            let base = type_name.rsplit("::").next().unwrap_or(type_name);
533            format!("{}Schema", base)
534        }
535    }
536}
537
538/// Extract the first generic argument from a type string.
539///
540/// `"Result<Json<Planet>, E>"` → `Some("Json<Planet>")`
541fn extract_first_generic_arg_string(type_str: &str) -> Option<String> {
542    let start = type_str.find('<')? + 1;
543    let mut depth = 0;
544    let mut end = start;
545
546    for (i, ch) in type_str[start..].char_indices() {
547        match ch {
548            '<' => depth += 1,
549            '>' if depth == 0 => {
550                end = start + i;
551                break;
552            }
553            '>' => depth -= 1,
554            ',' if depth == 0 => {
555                end = start + i;
556                break;
557            }
558            _ => {}
559        }
560    }
561
562    if end > start {
563        Some(type_str[start..end].to_string())
564    } else {
565        None
566    }
567}
568
569/// Convert type name to schema constant name: `"Planet"` → `"PlanetSchema"`
570pub fn to_schema_name(rust_type: &str) -> String {
571    format!("{}Schema", base_type_name(rust_type))
572}
573
574/// Extract the base type name, stripping all wrappers.
575///
576/// `"Result<Json<Vec<Planet>>, E>"` → `"Planet"`
577pub fn base_type_name(rust_type: &str) -> String {
578    let mut base = rust_type.trim();
579
580    if base.starts_with("Result<")
581        && let Some(inner) = extract_first_generic_arg_string(base)
582    {
583        base = Box::leak(inner.into_boxed_str());
584    }
585    if base.starts_with("Json<") && base.ends_with('>') {
586        base = &base[5..base.len() - 1];
587    }
588    if base.starts_with("Vec<") && base.ends_with('>') {
589        base = &base[4..base.len() - 1];
590    }
591    if base.starts_with("Option<") && base.ends_with('>') {
592        base = &base[7..base.len() - 1];
593    }
594
595    base.rsplit("::").next().unwrap_or(base).to_string()
596}
597
598#[cfg(test)]
599mod runtime_conversion_tests {
600    use super::*;
601
602    #[test]
603    fn json_planet() {
604        assert_eq!(rust_type_to_ts_schema("Json<Planet>"), "PlanetSchema");
605    }
606
607    #[test]
608    fn json_vec_planet() {
609        assert_eq!(
610            rust_type_to_ts_schema("Json<Vec<Planet>>"),
611            "z.array(PlanetSchema)"
612        );
613    }
614
615    #[test]
616    fn result_json_planet() {
617        assert_eq!(
618            rust_type_to_ts_schema("Result<Json<Planet>, StatusCode>"),
619            "PlanetSchema"
620        );
621    }
622
623    #[test]
624    fn json_string() {
625        assert_eq!(rust_type_to_ts_schema("Json<String>"), "z.string()");
626    }
627
628    #[test]
629    fn unit_type() {
630        assert_eq!(rust_type_to_ts_schema("()"), "");
631    }
632
633    #[test]
634    fn serde_json_value() {
635        assert_eq!(rust_type_to_ts_schema("Json<serde_json::Value>"), "z.any()");
636    }
637
638    #[test]
639    fn schema_name_simple() {
640        assert_eq!(to_schema_name("Planet"), "PlanetSchema");
641    }
642
643    #[test]
644    fn schema_name_vec() {
645        assert_eq!(to_schema_name("Vec<Planet>"), "PlanetSchema");
646    }
647
648    #[test]
649    fn base_type_unwraps_wrappers() {
650        assert_eq!(base_type_name("Result<Json<Vec<Planet>>, E>"), "Planet");
651        assert_eq!(base_type_name("Json<Planet>"), "Planet");
652        assert_eq!(base_type_name("Vec<Planet>"), "Planet");
653        assert_eq!(base_type_name("Option<Planet>"), "Planet");
654    }
655
656    #[test]
657    fn base_type_strips_module_path() {
658        assert_eq!(base_type_name("models::Planet"), "Planet");
659        assert_eq!(base_type_name("crate::domain::Planet"), "Planet");
660    }
661}