1use 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
19pub 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
44fn 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 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
104fn 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
149fn 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
223fn 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
258pub fn rust_type_to_zod(ty: &syn::Type, attrs: &ZodAttrs) -> String {
267 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 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 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" => "z.uuid()".to_string(),
306 "DateTime" => "z.iso.datetime({ offset: true })".to_string(),
308 "Value" => "z.any()".to_string(),
310 other => format!("{}Schema", other),
312 };
313 }
314
315 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
325fn 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
400fn 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
412fn innermost_custom_name(ty: &syn::Type) -> Option<String> {
415 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 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
455pub 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 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 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
511fn 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
538fn 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
569pub fn to_schema_name(rust_type: &str) -> String {
571 format!("{}Schema", base_type_name(rust_type))
572}
573
574pub 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}