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 "Value" => "z.any()".to_string(),
306 other => format!("{}Schema", other),
308 };
309 }
310
311 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
321fn 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
396fn 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
408fn innermost_custom_name(ty: &syn::Type) -> Option<String> {
411 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 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
451pub 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 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 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
507fn 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
532fn 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
563pub fn to_schema_name(rust_type: &str) -> String {
565 format!("{}Schema", base_type_name(rust_type))
566}
567
568pub 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}