1use std::collections::{BTreeMap, BTreeSet};
4use std::fs;
5use std::path::Path;
6
7use crate::casing::Casing;
8use crate::error::{Diagnostic, DiagnosticKind, Error, SourceLocation};
9use crate::expr::{CodecExpr, generate_import_block};
10use crate::registry::{ExternalType, Registry, WithWrapper};
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub enum OnUnknown {
15 #[default]
17 Error,
18 SkipContainingType,
21}
22
23#[derive(Debug, Clone)]
25pub enum EnumVariant {
26 Unit(String),
28 Newtype(String, CodecExpr),
30 Tuple(String, Vec<CodecExpr>),
33 Struct(String, Vec<(String, CodecExpr)>),
35}
36
37impl EnumVariant {
38 pub fn name(&self) -> &str {
40 match self {
41 EnumVariant::Unit(name)
42 | EnumVariant::Newtype(name, _)
43 | EnumVariant::Tuple(name, _)
44 | EnumVariant::Struct(name, _) => name,
45 }
46 }
47}
48
49#[derive(Debug, Clone)]
51pub(crate) enum TypeKind {
52 Struct(Vec<(String, CodecExpr)>),
53 Enum(Vec<EnumVariant>),
54 Alias(CodecExpr),
55}
56
57#[derive(Debug, Clone)]
59struct FormatSpec {
60 endian: String,
61 pointer_width: u32,
62 aligned: bool,
63}
64
65impl FormatSpec {
66 fn is_default(&self) -> bool {
67 self.endian == "little" && self.pointer_width == 32 && self.aligned
68 }
69
70 fn options(&self) -> String {
72 let mut entries = Vec::new();
73 if self.endian != "little" {
74 entries.push(format!("endian: '{}'", self.endian));
75 }
76 if self.pointer_width != 32 {
77 entries.push(format!("pointerWidth: {}", self.pointer_width));
78 }
79 if !self.aligned {
80 entries.push("aligned: false".to_string());
81 }
82 entries.join(", ")
83 }
84}
85
86#[derive(Debug)]
101pub struct CodeGenerator {
102 pub(crate) types: BTreeMap<String, TypeKind>,
104 pub(crate) failed: BTreeMap<String, Vec<Diagnostic>>,
106 pub(crate) add_diagnostics: Vec<Diagnostic>,
108 overrides: BTreeMap<String, String>,
110 header: Option<String>,
111 allow_typescript_syntax: bool,
112 pub(crate) on_unknown: OnUnknown,
113 pub(crate) marker_paths: BTreeSet<String>,
115 pub(crate) registry: Registry,
116 format: Option<FormatSpec>,
117 direction: Direction,
118 jit: bool,
119 field_casing: Casing,
120 variant_casing: Casing,
121}
122
123#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
128pub enum Direction {
129 #[default]
131 Full,
132 Decode,
134 Encode,
136}
137
138impl Direction {
139 fn suffix(self) -> Option<&'static str> {
140 match self {
141 Direction::Full => None,
142 Direction::Decode => Some("/decode"),
143 Direction::Encode => Some("/encode"),
144 }
145 }
146
147 fn jit_entry(self) -> (&'static str, &'static str) {
149 match self {
150 Direction::Full => ("rkyv-js/jit", "compileCodec"),
151 Direction::Decode => ("rkyv-js/jit/decode", "compileDecoder"),
152 Direction::Encode => ("rkyv-js/jit/encode", "compileEncoder"),
153 }
154 }
155
156 pub(crate) fn rewrite_import_block(self, block: &str) -> String {
160 let Some(suffix) = self.suffix() else {
161 return block.to_string();
162 };
163 let mut out = String::with_capacity(block.len() + 64);
164 for line in block.lines() {
165 if let Some(spec_start) = line.rfind(" from '").map(|i| i + " from '".len())
166 && let Some(len) = line[spec_start..].find('\'')
167 {
168 let spec = &line[spec_start..spec_start + len];
169 if spec == "rkyv-js" || spec.starts_with("rkyv-js/lib/") {
170 out.push_str(&line[..spec_start + len]);
171 out.push_str(suffix);
172 out.push_str(&line[spec_start + len..]);
173 out.push('\n');
174 continue;
175 }
176 }
177 out.push_str(line);
178 out.push('\n');
179 }
180 out
181 }
182}
183
184impl Default for CodeGenerator {
185 fn default() -> Self {
186 Self::new()
187 }
188}
189
190impl CodeGenerator {
191 pub fn new() -> Self {
193 Self {
194 types: BTreeMap::new(),
195 failed: BTreeMap::new(),
196 add_diagnostics: Vec::new(),
197 overrides: BTreeMap::new(),
198 header: None,
199 allow_typescript_syntax: true,
200 on_unknown: OnUnknown::Error,
201 marker_paths: BTreeSet::from(["rkyv::Archive".to_string()]),
202 registry: Registry::with_builtins(),
203 format: None,
204 direction: Direction::Full,
205 jit: false,
206 field_casing: Casing::Preserve,
207 variant_casing: Casing::Preserve,
208 }
209 }
210
211 pub fn set_header(&mut self, header: impl Into<String>) -> &mut Self {
213 self.header = Some(header.into());
214 self
215 }
216
217 pub fn set_direction(&mut self, direction: Direction) -> &mut Self {
223 self.direction = direction;
224 self
225 }
226
227 pub fn set_jit(&mut self, enabled: bool) -> &mut Self {
243 self.jit = enabled;
244 self
245 }
246
247 pub fn set_field_casing(&mut self, casing: Casing) -> &mut Self {
273 self.field_casing = casing;
274 self
275 }
276
277 pub fn set_variant_casing(&mut self, casing: Casing) -> &mut Self {
289 self.variant_casing = casing;
290 self
291 }
292
293 pub fn allow_typescript_syntax(&mut self, enabled: bool) -> &mut Self {
297 self.allow_typescript_syntax = enabled;
298 self
299 }
300
301 pub fn on_unknown_type(&mut self, mode: OnUnknown) -> &mut Self {
305 self.on_unknown = mode;
306 self
307 }
308
309 pub fn add_marker_path(&mut self, path: impl Into<String>) -> &mut Self {
312 self.marker_paths.insert(path.into());
313 self
314 }
315
316 pub fn register_external(
330 &mut self,
331 path: impl Into<String>,
332 external: ExternalType,
333 ) -> &mut Self {
334 self.registry.register_type(path, external);
335 self
336 }
337
338 pub fn register_with(&mut self, path: impl Into<String>, wrapper: WithWrapper) -> &mut Self {
340 self.registry.register_wrapper(path, wrapper);
341 self
342 }
343
344 pub fn unregister_external(&mut self, path: &str) -> &mut Self {
346 self.registry.unregister_type(path);
347 self
348 }
349
350 pub fn add_struct(
352 &mut self,
353 name: impl Into<String>,
354 fields: impl IntoIterator<Item = (impl Into<String>, CodecExpr)>,
355 ) -> &mut Self {
356 let fields = fields
357 .into_iter()
358 .map(|(field, expr)| (field.into(), expr))
359 .collect();
360 self.add_type(name.into(), TypeKind::Struct(fields), None);
361 self
362 }
363
364 pub fn add_enum(
366 &mut self,
367 name: impl Into<String>,
368 variants: impl IntoIterator<Item = EnumVariant>,
369 ) -> &mut Self {
370 let variants = variants.into_iter().collect();
371 self.add_type(name.into(), TypeKind::Enum(variants), None);
372 self
373 }
374
375 pub fn add_alias(&mut self, name: impl Into<String>, target: CodecExpr) -> &mut Self {
377 self.add_type(name.into(), TypeKind::Alias(target), None);
378 self
379 }
380
381 pub(crate) fn add_type(
383 &mut self,
384 name: String,
385 kind: TypeKind,
386 location: Option<SourceLocation>,
387 ) {
388 if self.is_known_type(&name) {
389 self.add_diagnostics.push(
390 Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
391 );
392 return;
393 }
394 self.types.insert(name, kind);
395 }
396
397 pub(crate) fn add_failed_type(
399 &mut self,
400 name: String,
401 diagnostics: Vec<Diagnostic>,
402 location: Option<SourceLocation>,
403 ) {
404 if self.is_known_type(&name) {
405 self.add_diagnostics.push(
406 Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
407 );
408 return;
409 }
410 self.failed.insert(name, diagnostics);
411 }
412
413 fn is_known_type(&self, name: &str) -> bool {
414 self.types.contains_key(name) || self.failed.contains_key(name)
415 }
416
417 pub fn set_archived_name(
422 &mut self,
423 type_name: impl Into<String>,
424 archived_name: impl Into<String>,
425 ) -> &mut Self {
426 self.overrides.insert(type_name.into(), archived_name.into());
427 self
428 }
429
430 pub fn archived_name_of(&self, type_name: &str) -> Option<String> {
433 if !self.is_known_type(type_name) {
434 return None;
435 }
436 Some(self.resolved_archived_name(type_name))
437 }
438
439 fn resolved_archived_name(&self, type_name: &str) -> String {
440 self.overrides
441 .get(type_name)
442 .cloned()
443 .unwrap_or_else(|| format!("Archived{type_name}"))
444 }
445
446 pub fn set_format(&mut self, endian: &str, pointer_width: u32, aligned: bool) -> &mut Self {
452 self.format = Some(FormatSpec {
453 endian: endian.to_string(),
454 pointer_width,
455 aligned,
456 });
457 self
458 }
459
460 fn nondefault_format(&self) -> Option<&FormatSpec> {
462 self.format.as_ref().filter(|spec| !spec.is_default())
463 }
464
465 fn exprs_with_context<'a>(
468 type_name: &str,
469 kind: &'a TypeKind,
470 ) -> Vec<(String, &'a CodecExpr)> {
471 match kind {
472 TypeKind::Struct(fields) => fields
473 .iter()
474 .map(|(field, expr)| (format!("{type_name}.{field}"), expr))
475 .collect(),
476 TypeKind::Enum(variants) => {
477 let mut out = Vec::new();
478 for variant in variants {
479 match variant {
480 EnumVariant::Unit(_) => {}
481 EnumVariant::Newtype(vname, expr) => {
482 out.push((format!("{type_name}::{vname}"), expr));
483 }
484 EnumVariant::Tuple(vname, exprs) => {
485 for (i, expr) in exprs.iter().enumerate() {
486 out.push((format!("{type_name}::{vname}.{i}"), expr));
487 }
488 }
489 EnumVariant::Struct(vname, fields) => {
490 for (field, expr) in fields {
491 out.push((format!("{type_name}::{vname}.{field}"), expr));
492 }
493 }
494 }
495 }
496 out
497 }
498 TypeKind::Alias(expr) => vec![(type_name.to_string(), expr)],
499 }
500 }
501
502 fn casing_collisions(
508 context: &str,
509 names: impl IntoIterator<Item = String>,
510 casing: Casing,
511 ) -> Vec<Diagnostic> {
512 if casing == Casing::Preserve {
513 return Vec::new();
514 }
515 let mut by_emitted: BTreeMap<String, Vec<String>> = BTreeMap::new();
516 for name in names {
517 by_emitted.entry(casing.apply(&name)).or_default().push(name);
518 }
519 by_emitted
520 .into_iter()
521 .filter(|(_, originals)| originals.len() > 1)
522 .map(|(emitted, originals)| {
523 Diagnostic::new(DiagnosticKind::NameCollision { emitted, originals })
524 .referenced_by(context.to_string())
525 })
526 .collect()
527 }
528
529 fn casing_diagnostics(&self, emitted: &BTreeMap<&String, &TypeKind>) -> Vec<Diagnostic> {
531 let mut diagnostics = Vec::new();
532 for (name, kind) in emitted {
533 match kind {
534 TypeKind::Struct(fields) => {
535 diagnostics.extend(Self::casing_collisions(
536 name,
537 fields.iter().map(|(field, _)| field.clone()),
538 self.field_casing,
539 ));
540 }
541 TypeKind::Enum(variants) => {
542 diagnostics.extend(Self::casing_collisions(
543 name,
544 variants.iter().map(|variant| variant.name().to_string()),
545 self.variant_casing,
546 ));
547 for variant in variants.iter() {
548 if let EnumVariant::Struct(vname, fields) = variant {
549 diagnostics.extend(Self::casing_collisions(
550 &format!("{name}::{vname}"),
551 fields.iter().map(|(field, _)| field.clone()),
552 self.field_casing,
553 ));
554 }
555 }
556 }
557 TypeKind::Alias(_) => {}
558 }
559 }
560 diagnostics
561 }
562
563 pub fn generate(&self) -> Result<String, Error> {
567 let mut diagnostics: Vec<Diagnostic> = self.add_diagnostics.clone();
568
569 for target in self.overrides.keys() {
571 if !self.is_known_type(target) {
572 diagnostics.push(Diagnostic::new(DiagnosticKind::UnknownRenameTarget {
573 type_name: target.clone(),
574 }));
575 }
576 }
577
578 let mut skipped: BTreeSet<String> = BTreeSet::new();
580 match self.on_unknown {
581 OnUnknown::Error => {
582 for failure_diagnostics in self.failed.values() {
583 diagnostics.extend(failure_diagnostics.iter().cloned());
584 }
585 }
586 OnUnknown::SkipContainingType => {
587 for (name, failure_diagnostics) in &self.failed {
588 skipped.insert(name.clone());
589 for diagnostic in failure_diagnostics {
590 eprintln!(
591 "cargo:warning=rkyv-js-codegen: skipping `{name}`: {diagnostic}"
592 );
593 }
594 }
595 }
596 }
597
598 match self.on_unknown {
600 OnUnknown::Error => {
601 for (name, kind) in &self.types {
602 for (context, expr) in Self::exprs_with_context(name, kind) {
603 let mut refs = BTreeSet::new();
604 expr.collect_type_refs(&mut refs);
605 for reference in refs {
606 if !self.is_known_type(&reference) {
607 diagnostics.push(
608 Diagnostic::new(DiagnosticKind::UnresolvedTypeRef {
609 name: reference,
610 })
611 .referenced_by(context.clone()),
612 );
613 }
614 }
615 }
616 }
617 }
618 OnUnknown::SkipContainingType => {
619 loop {
621 let mut newly_skipped = Vec::new();
622 for (name, kind) in &self.types {
623 if skipped.contains(name) {
624 continue;
625 }
626 let broken = Self::exprs_with_context(name, kind).iter().any(
627 |(_, expr)| {
628 let mut refs = BTreeSet::new();
629 expr.collect_type_refs(&mut refs);
630 refs.iter().any(|reference| {
631 skipped.contains(reference)
632 || !self.types.contains_key(reference)
633 })
634 },
635 );
636 if broken {
637 newly_skipped.push(name.clone());
638 }
639 }
640 if newly_skipped.is_empty() {
641 break;
642 }
643 for name in newly_skipped {
644 eprintln!(
645 "cargo:warning=rkyv-js-codegen: skipping `{name}`: it references \
646 a type that was omitted or never added"
647 );
648 skipped.insert(name);
649 }
650 }
651 }
652 }
653
654 let emitted: BTreeMap<&String, &TypeKind> = self
656 .types
657 .iter()
658 .filter(|(name, _)| !skipped.contains(*name))
659 .collect();
660
661 diagnostics.extend(self.casing_diagnostics(&emitted));
662
663 let (jit_module, jit_fn) = self.direction.jit_entry();
665 let jit_import = CodecExpr::import_from(jit_module, jit_fn);
666 let mut all_exprs: Vec<&CodecExpr> = emitted
667 .iter()
668 .flat_map(|(name, kind)| Self::exprs_with_context(name, kind))
669 .map(|(_, expr)| expr)
670 .collect();
671 if self.jit && !emitted.is_empty() {
672 all_exprs.push(&jit_import);
675 }
676 let import_block = match generate_import_block(all_exprs.iter().copied()) {
677 Ok(block) => self.direction.rewrite_import_block(&block),
678 Err(conflicts) => {
679 diagnostics.extend(conflicts.into_iter().map(Diagnostic::new));
680 String::new()
681 }
682 };
683
684 if !diagnostics.is_empty() {
685 return Err(Error::Codegen(diagnostics));
686 }
687
688 let order = Self::topological_sort(&emitted);
691
692 let archived_names: BTreeMap<String, String> = emitted
693 .keys()
694 .map(|name| ((*name).clone(), self.resolved_archived_name(name)))
695 .collect();
696
697 let codec_names: BTreeMap<String, String> = if self.jit {
700 archived_names
701 .iter()
702 .map(|(name, archived)| (name.clone(), format!("{archived}$")))
703 .collect()
704 } else {
705 archived_names.clone()
706 };
707
708 let mut blocks: Vec<String> = Vec::new();
710
711 let header = self
712 .header
713 .as_deref()
714 .unwrap_or("Auto-generated by rkyv-js-codegen\nDO NOT EDIT MANUALLY");
715 let mut header_block = String::from("/**\n");
716 for line in header.lines() {
717 if line.is_empty() {
718 header_block.push_str(" *\n");
719 } else {
720 header_block.push_str(" * ");
721 header_block.push_str(line);
722 header_block.push('\n');
723 }
724 }
725 header_block.push_str(" */");
726 blocks.push(header_block);
727
728 blocks.push(import_block.trim_end().to_string());
729
730 if let Some(spec) = self.nondefault_format() {
731 blocks.push(format!("const FORMAT = r.format({{ {} }});", spec.options()));
732 }
733
734 for name in &order {
735 let kind = emitted.get(name).expect("ordered names come from emitted");
736 blocks.push(self.emit_type(name, kind, &archived_names, &codec_names));
737 }
738
739 Ok(blocks.join("\n\n") + "\n")
740 }
741
742 fn topological_sort(emitted: &BTreeMap<&String, &TypeKind>) -> Vec<String> {
743 let mut deps: BTreeMap<&str, BTreeSet<String>> = BTreeMap::new();
744 for (name, kind) in emitted {
745 let mut refs = BTreeSet::new();
746 for (_, expr) in Self::exprs_with_context(name, kind) {
747 expr.collect_type_refs(&mut refs);
748 }
749 refs.retain(|reference| {
750 emitted.contains_key(reference) && reference != name.as_str()
751 });
752 deps.insert(name.as_str(), refs);
753 }
754
755 let mut dependents: BTreeMap<&str, Vec<&str>> = BTreeMap::new();
756 let mut in_degree: BTreeMap<&str, usize> = BTreeMap::new();
757 for (name, type_deps) in &deps {
758 in_degree.insert(name, type_deps.len());
759 for dep in type_deps {
760 dependents.entry(dep.as_str()).or_default().push(name);
761 }
762 }
763
764 let mut ready: BTreeSet<&str> = in_degree
765 .iter()
766 .filter(|(_, degree)| **degree == 0)
767 .map(|(name, _)| *name)
768 .collect();
769 let mut order: Vec<String> = Vec::new();
770 let mut done: BTreeSet<&str> = BTreeSet::new();
771
772 while let Some(name) = ready.pop_first() {
773 order.push(name.to_string());
774 done.insert(name);
775 if let Some(children) = dependents.get(name) {
776 for child in children {
777 let degree = in_degree.get_mut(child).unwrap();
778 *degree -= 1;
779 if *degree == 0 {
780 ready.insert(child);
781 }
782 }
783 }
784 }
785
786 for name in deps.keys() {
789 if !done.contains(name) {
790 order.push((*name).to_string());
791 }
792 }
793
794 order
795 }
796
797 fn emit_type(
798 &self,
799 name: &str,
800 kind: &TypeKind,
801 archived_names: &BTreeMap<String, String>,
802 codec_names: &BTreeMap<String, String>,
803 ) -> String {
804 let archived = archived_names
805 .get(name)
806 .expect("emitted types have archived names")
807 .clone();
808 let render = |expr: &CodecExpr| -> String {
809 expr.render(codec_names)
810 .expect("type references are validated before emission")
811 };
812
813 let codec_expr = match kind {
814 TypeKind::Struct(fields) => {
815 if fields.is_empty() {
816 "r.struct({})".to_string()
817 } else {
818 let mut body = String::from("r.struct({\n");
819 for (field, expr) in fields {
820 body.push_str(&format!(
821 " {}: {},\n",
822 self.field_casing.apply(field),
823 render(expr)
824 ));
825 }
826 body.push_str("})");
827 body
828 }
829 }
830 TypeKind::Enum(variants) => {
831 if variants.is_empty() {
832 "r.taggedEnum({})".to_string()
833 } else {
834 let mut body = String::from("r.taggedEnum({\n");
835 for variant in variants {
836 let value = match variant {
837 EnumVariant::Unit(_) => "null".to_string(),
838 EnumVariant::Newtype(_, expr) => render(expr),
839 EnumVariant::Tuple(_, exprs) => {
840 render(&CodecExpr::array(exprs.iter().cloned()))
841 }
842 EnumVariant::Struct(_, fields) => {
843 let record = CodecExpr::object(fields.iter().map(
844 |(field, expr)| {
845 (self.field_casing.apply(field), expr.clone())
846 },
847 ));
848 render(&record)
849 }
850 };
851 body.push_str(&format!(
852 " {}: {},\n",
853 self.variant_casing.apply(variant.name()),
854 value
855 ));
856 }
857 body.push_str("})");
858 body
859 }
860 }
861 TypeKind::Alias(expr) => render(expr),
862 };
863
864 let codec_expr = match self.nondefault_format() {
865 Some(_) => format!("r.withFormat({codec_expr}, FORMAT)"),
866 None => codec_expr,
867 };
868
869 let mut block = if self.jit {
870 let jit_fn = self.direction.jit_entry().1;
873 format!(
874 "const {archived}$ = {codec_expr};\n\n\
875 export const {archived} = {jit_fn}({archived}$);"
876 )
877 } else {
878 format!("export const {archived} = {codec_expr};")
879 };
880 if self.allow_typescript_syntax {
881 block.push_str(&format!(
882 "\n\nexport type {name} = r.Infer<typeof {archived}>;"
883 ));
884 }
885 block
886 }
887
888 pub fn write_to_file(&self, path: impl AsRef<Path>) -> Result<(), Error> {
890 let code = self.generate()?;
891 fs::write(path, code)?;
892 Ok(())
893 }
894}
895
896#[cfg(test)]
897mod tests {
898 use super::*;
899 use crate::expr::codec;
900
901 fn diagnostics(error: Error) -> Vec<Diagnostic> {
902 match error {
903 Error::Codegen(diagnostics) => diagnostics,
904 other => panic!("expected Error::Codegen, got {other:?}"),
905 }
906 }
907
908 #[test]
909 fn struct_emission_snapshot() {
910 let mut generator = CodeGenerator::new();
911 generator.add_struct("Point", [("x", codec::f64()), ("y", codec::f64())]);
912 let code = generator.generate().unwrap();
913 assert_eq!(
914 code,
915 "/**\n\
916 \x20* Auto-generated by rkyv-js-codegen\n\
917 \x20* DO NOT EDIT MANUALLY\n\
918 \x20*/\n\
919 \n\
920 import * as r from 'rkyv-js';\n\
921 \n\
922 export const ArchivedPoint = r.struct({\n\
923 \x20 x: r.f64,\n\
924 \x20 y: r.f64,\n\
925 });\n\
926 \n\
927 export type Point = r.Infer<typeof ArchivedPoint>;\n"
928 );
929 }
930
931 #[test]
932 fn enum_emission_snapshot() {
933 let mut generator = CodeGenerator::new();
934 generator.add_enum(
935 "MixedAlign",
936 [
937 EnumVariant::Struct(
938 "V".to_string(),
939 vec![("a".to_string(), codec::u8()), ("b".to_string(), codec::u32())],
940 ),
941 EnumVariant::Newtype("X".to_string(), codec::u64()),
942 EnumVariant::Tuple("Color".to_string(), vec![codec::u8(), codec::u8()]),
943 EnumVariant::Unit("Y".to_string()),
944 ],
945 );
946 let code = generator.generate().unwrap();
947 assert!(code.contains(
948 "export const ArchivedMixedAlign = r.taggedEnum({\n\
949 \x20 V: { a: r.u8, b: r.u32 },\n\
950 \x20 X: r.u64,\n\
951 \x20 Color: [r.u8, r.u8],\n\
952 \x20 Y: null,\n\
953 });"
954 ));
955 assert!(code.contains("export type MixedAlign = r.Infer<typeof ArchivedMixedAlign>;"));
956 }
957
958 #[test]
959 fn alias_emission_snapshot() {
960 let mut generator = CodeGenerator::new();
961 generator.add_alias("UserId", codec::u32());
962 let code = generator.generate().unwrap();
963 assert!(code.contains("export const ArchivedUserId = r.u32;"));
964 assert!(code.contains("export type UserId = r.Infer<typeof ArchivedUserId>;"));
965 }
966
967 #[test]
968 fn imports_are_collected_and_deduped() {
969 let mut generator = CodeGenerator::new();
970 generator.add_struct(
971 "A",
972 [
973 (
974 "m",
975 CodecExpr::call(
976 CodecExpr::import_from("rkyv-js/lib/hashmap", "hashMap"),
977 [codec::string(), codec::u32()],
978 ),
979 ),
980 (
981 "s",
982 CodecExpr::call(
983 CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
984 [codec::string()],
985 ),
986 ),
987 ],
988 );
989 generator.add_struct(
990 "B",
991 [(
992 "s2",
993 CodecExpr::call(
994 CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
995 [codec::u32()],
996 ),
997 )],
998 );
999 let code = generator.generate().unwrap();
1000 assert!(code.contains("import { hashMap, hashSet } from 'rkyv-js/lib/hashmap';"));
1001 assert_eq!(code.matches("hashSet }").count(), 1);
1002 }
1003
1004 #[test]
1005 fn import_conflict_is_reported() {
1006 let mut generator = CodeGenerator::new();
1007 generator.add_struct("A", [("x", CodecExpr::import_from("pkg-a", "codec"))]);
1008 generator.add_struct("B", [("y", CodecExpr::import_from("pkg-b", "codec"))]);
1009 let errors = diagnostics(generator.generate().unwrap_err());
1010 assert!(errors.iter().any(|diagnostic| matches!(
1011 &diagnostic.kind,
1012 DiagnosticKind::ImportConflict { export, .. } if export == "codec"
1013 )));
1014 }
1015
1016 #[test]
1017 fn topo_sort_handles_forward_references() {
1018 let mut generator = CodeGenerator::new();
1019 generator.add_struct("AOuter", [("inner", codec::named("Inner"))]);
1021 generator.add_struct("Inner", [("value", codec::u32())]);
1022 let code = generator.generate().unwrap();
1023 let inner_pos = code.find("export const ArchivedInner").unwrap();
1024 let outer_pos = code.find("export const ArchivedAOuter").unwrap();
1025 assert!(inner_pos < outer_pos, "dependency must be emitted first");
1026 assert!(code.contains("inner: ArchivedInner,"));
1027 }
1028
1029 #[test]
1030 fn unresolved_type_ref_reports_referrer() {
1031 let mut generator = CodeGenerator::new();
1032 generator.add_struct("Outer", [("inner", codec::named("Missing"))]);
1033 let errors = diagnostics(generator.generate().unwrap_err());
1034 assert_eq!(errors.len(), 1);
1035 assert!(matches!(
1036 &errors[0].kind,
1037 DiagnosticKind::UnresolvedTypeRef { name } if name == "Missing"
1038 ));
1039 assert_eq!(errors[0].referenced_by.as_deref(), Some("Outer.inner"));
1040 }
1041
1042 #[test]
1043 fn duplicate_type_is_reported_at_generate() {
1044 let mut generator = CodeGenerator::new();
1045 generator.add_struct("Point", [("x", codec::f64())]);
1046 generator.add_struct("Point", [("y", codec::f64())]);
1047 let errors = diagnostics(generator.generate().unwrap_err());
1048 assert!(errors.iter().any(|diagnostic| matches!(
1049 &diagnostic.kind,
1050 DiagnosticKind::DuplicateType { name } if name == "Point"
1051 )));
1052 }
1053
1054 #[test]
1055 fn set_archived_name_is_order_independent() {
1056 let mut generator = CodeGenerator::new();
1058 generator.set_archived_name("Foo", "MyFoo");
1059 generator.add_struct("Foo", [("x", codec::u32())]);
1060 let code = generator.generate().unwrap();
1061 assert!(code.contains("export const MyFoo = r.struct({"));
1062 assert!(code.contains("export type Foo = r.Infer<typeof MyFoo>;"));
1063 assert!(!code.contains("ArchivedFoo"));
1064
1065 let mut generator = CodeGenerator::new();
1067 generator.add_struct("Foo", [("x", codec::u32())]);
1068 generator.set_archived_name("Foo", "MyFoo");
1069 let code = generator.generate().unwrap();
1070 assert!(code.contains("export const MyFoo = r.struct({"));
1071 }
1072
1073 #[test]
1074 fn archived_rename_applies_to_cross_references() {
1075 let mut generator = CodeGenerator::new();
1076 generator.set_archived_name("Inner", "CustomInner");
1077 generator.add_struct("Inner", [("value", codec::u32())]);
1078 generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
1079 let code = generator.generate().unwrap();
1080 assert!(code.contains("export const CustomInner = r.struct({"));
1081 assert!(code.contains("inner: CustomInner,"));
1082 assert!(!code.contains("ArchivedInner"));
1083 }
1084
1085 #[test]
1086 fn unknown_rename_target_is_a_diagnostic() {
1087 let mut generator = CodeGenerator::new();
1088 generator.add_struct("Foo", [("x", codec::u32())]);
1089 generator.set_archived_name("Nope", "MyNope");
1090 let errors = diagnostics(generator.generate().unwrap_err());
1091 assert!(errors.iter().any(|diagnostic| matches!(
1092 &diagnostic.kind,
1093 DiagnosticKind::UnknownRenameTarget { type_name } if type_name == "Nope"
1094 )));
1095 }
1096
1097 #[test]
1098 fn archived_name_of_accessor() {
1099 let mut generator = CodeGenerator::new();
1100 generator.add_struct("Foo", [("x", codec::u32())]);
1101 assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("ArchivedFoo"));
1102 generator.set_archived_name("Foo", "MyFoo");
1103 assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("MyFoo"));
1104 assert_eq!(generator.archived_name_of("Bar"), None);
1105 }
1106
1107 #[test]
1108 fn field_casing_defaults_to_preserve() {
1109 let mut generator = CodeGenerator::new();
1110 generator.add_struct("Event", [("created_at", codec::u64())]);
1111 let code = generator.generate().unwrap();
1112 assert!(code.contains("created_at: r.u64,"));
1113 }
1114
1115 #[test]
1116 fn field_casing_camel_rewrites_struct_fields() {
1117 let mut generator = CodeGenerator::new();
1118 generator.set_field_casing(Casing::Camel);
1119 generator.add_struct(
1120 "Event",
1121 [
1122 ("created_at", codec::u64()),
1123 ("HTTP_status", codec::u16()),
1124 ("id", codec::u32()),
1125 ],
1126 );
1127 let code = generator.generate().unwrap();
1128 assert!(code.contains("createdAt: r.u64,"));
1129 assert!(code.contains("httpStatus: r.u16,"));
1130 assert!(code.contains("id: r.u32,"));
1131 let created = code.find("createdAt").unwrap();
1133 let status = code.find("httpStatus").unwrap();
1134 assert!(created < status);
1135 }
1136
1137 #[test]
1138 fn field_casing_applies_to_enum_struct_variants() {
1139 let mut generator = CodeGenerator::new();
1140 generator.set_field_casing(Casing::Camel);
1141 generator.add_enum(
1142 "Message",
1143 [
1144 EnumVariant::Struct(
1145 "Text".to_string(),
1146 vec![
1147 ("sent_at".to_string(), codec::u64()),
1148 ("body_text".to_string(), codec::string()),
1149 ],
1150 ),
1151 EnumVariant::Unit("Ping".to_string()),
1152 ],
1153 );
1154 let code = generator.generate().unwrap();
1155 assert!(code.contains("Text: { sentAt: r.u64, bodyText: r.string },"));
1156 assert!(code.contains("Ping: null,"));
1158 }
1159
1160 #[test]
1161 fn variant_casing_is_independent_of_field_casing() {
1162 let mut generator = CodeGenerator::new();
1163 generator
1164 .set_field_casing(Casing::Camel)
1165 .set_variant_casing(Casing::Snake);
1166 generator.add_enum(
1167 "Message",
1168 [
1169 EnumVariant::Struct(
1170 "PlainText".to_string(),
1171 vec![("sent_at".to_string(), codec::u64())],
1172 ),
1173 EnumVariant::Newtype("BinaryBlob".to_string(), codec::string()),
1174 ],
1175 );
1176 let code = generator.generate().unwrap();
1177 assert!(code.contains("plain_text: { sentAt: r.u64 },"));
1178 assert!(code.contains("binary_blob: r.string,"));
1179 }
1180
1181 #[test]
1182 fn casing_leaves_type_and_export_names_alone() {
1183 let mut generator = CodeGenerator::new();
1184 generator.set_field_casing(Casing::Camel);
1185 generator.add_struct("HttpEvent", [("created_at", codec::u64())]);
1186 let code = generator.generate().unwrap();
1187 assert!(code.contains("export const ArchivedHttpEvent = r.struct({"));
1188 assert!(code.contains("export type HttpEvent = r.Infer<typeof ArchivedHttpEvent>;"));
1189 }
1190
1191 #[test]
1192 fn casing_collision_is_reported() {
1193 let mut generator = CodeGenerator::new();
1194 generator.set_field_casing(Casing::Camel);
1195 generator.add_struct(
1196 "Event",
1197 [("foo_bar", codec::u32()), ("fooBar", codec::u32())],
1198 );
1199 let errors = diagnostics(generator.generate().unwrap_err());
1200 assert_eq!(errors.len(), 1);
1201 assert!(matches!(
1202 &errors[0].kind,
1203 DiagnosticKind::NameCollision { emitted, originals }
1204 if emitted == "fooBar" && originals.len() == 2
1205 ));
1206 assert_eq!(errors[0].referenced_by.as_deref(), Some("Event"));
1207 }
1208
1209 #[test]
1210 fn casing_collision_in_a_struct_variant_names_the_variant() {
1211 let mut generator = CodeGenerator::new();
1212 generator.set_field_casing(Casing::Camel);
1213 generator.add_enum(
1214 "Message",
1215 [EnumVariant::Struct(
1216 "Text".to_string(),
1217 vec![
1218 ("sent_at".to_string(), codec::u64()),
1219 ("sentAt".to_string(), codec::u64()),
1220 ],
1221 )],
1222 );
1223 let errors = diagnostics(generator.generate().unwrap_err());
1224 assert_eq!(errors.len(), 1);
1225 assert_eq!(errors[0].referenced_by.as_deref(), Some("Message::Text"));
1226 }
1227
1228 #[test]
1229 fn variant_casing_collision_is_reported() {
1230 let mut generator = CodeGenerator::new();
1231 generator.set_variant_casing(Casing::Snake);
1232 generator.add_enum(
1233 "Message",
1234 [
1235 EnumVariant::Unit("PlainText".to_string()),
1236 EnumVariant::Unit("plain_text".to_string()),
1237 ],
1238 );
1239 let errors = diagnostics(generator.generate().unwrap_err());
1240 assert!(errors.iter().any(|diagnostic| matches!(
1241 &diagnostic.kind,
1242 DiagnosticKind::NameCollision { emitted, .. } if emitted == "plain_text"
1243 )));
1244 }
1245
1246 #[test]
1247 fn preserve_never_reports_a_collision() {
1248 let mut generator = CodeGenerator::new();
1250 generator.add_struct(
1251 "Event",
1252 [("foo_bar", codec::u32()), ("fooBar", codec::u32())],
1253 );
1254 let code = generator.generate().unwrap();
1255 assert!(code.contains("foo_bar: r.u32,"));
1256 assert!(code.contains("fooBar: r.u32,"));
1257 }
1258
1259 #[test]
1260 fn js_mode_omits_type_lines() {
1261 let mut generator = CodeGenerator::new();
1262 generator.allow_typescript_syntax(false);
1263 generator.add_struct("Point", [("x", codec::f64())]);
1264 generator.add_alias("UserId", codec::u32());
1265 let code = generator.generate().unwrap();
1266 assert!(code.contains("export const ArchivedPoint = r.struct({"));
1267 assert!(code.contains("export const ArchivedUserId = r.u32;"));
1268 assert!(!code.contains("export type"));
1269 assert!(!code.contains("r.Infer"));
1270 }
1271
1272 #[test]
1273 fn set_format_default_is_a_no_op() {
1274 let mut generator = CodeGenerator::new();
1275 generator.set_format("little", 32, true);
1276 generator.add_struct("Point", [("x", codec::f64())]);
1277 let code = generator.generate().unwrap();
1278 assert!(!code.contains("FORMAT"));
1279 assert!(!code.contains("withFormat"));
1280 }
1281
1282 #[test]
1283 fn set_format_nondefault_wraps_exports() {
1284 let mut generator = CodeGenerator::new();
1285 generator.set_format("big", 64, false);
1286 generator.add_struct("Point", [("x", codec::f64())]);
1287 generator.add_alias("UserId", codec::u32());
1288 let code = generator.generate().unwrap();
1289 assert!(code.contains(
1290 "const FORMAT = r.format({ endian: 'big', pointerWidth: 64, aligned: false });"
1291 ));
1292 assert!(code.contains("export const ArchivedPoint = r.withFormat(r.struct({\n"));
1293 assert!(code.contains("}), FORMAT);"));
1294 assert!(code.contains("export const ArchivedUserId = r.withFormat(r.u32, FORMAT);"));
1295 }
1296
1297 #[test]
1298 fn set_format_emits_only_nondefault_keys() {
1299 let mut generator = CodeGenerator::new();
1300 generator.set_format("little", 16, true);
1301 generator.add_struct("Point", [("x", codec::f64())]);
1302 let code = generator.generate().unwrap();
1303 assert!(code.contains("const FORMAT = r.format({ pointerWidth: 16 });"));
1304 }
1305
1306 #[test]
1307 fn custom_header_replaces_default() {
1308 let mut generator = CodeGenerator::new();
1309 generator.set_header("Custom header\nsecond line");
1310 generator.add_struct("Point", [("x", codec::f64())]);
1311 let code = generator.generate().unwrap();
1312 assert!(code.starts_with("/**\n * Custom header\n * second line\n */\n"));
1313 assert!(!code.contains("Auto-generated by rkyv-js-codegen"));
1314 }
1315
1316 #[test]
1317 fn set_direction_full_is_a_no_op() {
1318 let mut generator = CodeGenerator::new();
1319 generator.set_direction(Direction::Full);
1320 generator.add_struct("Point", [("x", codec::f64())]);
1321 let code = generator.generate().unwrap();
1322 assert!(code.contains("import * as r from 'rkyv-js';"));
1323 }
1324
1325 #[test]
1326 fn set_direction_rewrites_rkyv_specifiers_only() {
1327 let mut generator = CodeGenerator::new();
1328 generator.set_direction(Direction::Decode);
1329 generator.add_struct(
1330 "Event",
1331 [
1332 ("id", codec::u32()),
1333 (
1334 "tags",
1335 CodecExpr::call(
1336 CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
1337 [codec::string()],
1338 ),
1339 ),
1340 (
1341 "custom",
1342 CodecExpr::import_from("./my-codec.ts", "MyCodec"),
1343 ),
1344 ],
1345 );
1346 let code = generator.generate().unwrap();
1347 assert!(code.contains("import * as r from 'rkyv-js/decode';"));
1348 assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap/decode';"));
1349 assert!(code.contains("import { MyCodec } from './my-codec.ts';"));
1351 assert!(code.contains("export const ArchivedEvent = r.struct({"));
1353 assert!(code.contains("export type Event = r.Infer<typeof ArchivedEvent>;"));
1354 }
1355
1356 #[test]
1357 fn set_direction_encode_uses_encode_suffix() {
1358 let mut generator = CodeGenerator::new();
1359 generator.set_direction(Direction::Encode);
1360 generator.add_struct("Point", [("x", codec::f64())]);
1361 let code = generator.generate().unwrap();
1362 assert!(code.contains("import * as r from 'rkyv-js/encode';"));
1363 }
1364
1365 #[test]
1366 fn set_jit_wraps_exports() {
1367 let mut generator = CodeGenerator::new();
1368 generator.set_jit(true);
1369 generator.add_struct("Point", [("x", codec::f64())]);
1370 generator.add_alias("UserId", codec::u32());
1371 let code = generator.generate().unwrap();
1372 assert!(code.contains("import { compileCodec } from 'rkyv-js/jit';"));
1373 assert!(code.contains("const ArchivedPoint$ = r.struct({\n"));
1374 assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
1375 assert!(code.contains("const ArchivedUserId$ = r.u32;"));
1376 assert!(code.contains("export const ArchivedUserId = compileCodec(ArchivedUserId$);"));
1377 assert!(code.contains("export type Point = r.Infer<typeof ArchivedPoint>;"));
1379 }
1380
1381 #[test]
1382 fn set_jit_references_resolve_to_raw_codecs() {
1383 let mut generator = CodeGenerator::new();
1384 generator.set_jit(true);
1385 generator.add_struct("Inner", [("value", codec::u32())]);
1386 generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
1387 let code = generator.generate().unwrap();
1388 assert!(code.contains("inner: ArchivedInner$,"));
1391 assert!(code.contains("export const ArchivedInner = compileCodec(ArchivedInner$);"));
1392 assert!(code.contains("export const ArchivedOuter = compileCodec(ArchivedOuter$);"));
1393 }
1394
1395 #[test]
1396 fn set_jit_composes_with_format() {
1397 let mut generator = CodeGenerator::new();
1398 generator.set_jit(true);
1399 generator.set_format("big", 64, true);
1400 generator.add_struct("Point", [("x", codec::f64())]);
1401 let code = generator.generate().unwrap();
1402 assert!(code.contains("const FORMAT = r.format({ endian: 'big', pointerWidth: 64 });"));
1403 assert!(code.contains("const ArchivedPoint$ = r.withFormat(r.struct({\n"));
1405 assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
1406 }
1407
1408 #[test]
1409 fn set_jit_respects_archived_renames() {
1410 let mut generator = CodeGenerator::new();
1411 generator.set_jit(true);
1412 generator.set_archived_name("Inner", "CustomInner");
1413 generator.add_struct("Inner", [("value", codec::u32())]);
1414 generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
1415 let code = generator.generate().unwrap();
1416 assert!(code.contains("inner: CustomInner$,"));
1417 assert!(code.contains("export const CustomInner = compileCodec(CustomInner$);"));
1418 }
1419
1420 #[test]
1421 fn set_jit_decode_direction_uses_compile_decoder() {
1422 let mut generator = CodeGenerator::new();
1423 generator.set_jit(true);
1424 generator.set_direction(Direction::Decode);
1425 generator.add_struct(
1426 "Event",
1427 [
1428 ("id", codec::u32()),
1429 (
1430 "tags",
1431 CodecExpr::call(
1432 CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
1433 [codec::string()],
1434 ),
1435 ),
1436 ],
1437 );
1438 let code = generator.generate().unwrap();
1439 assert!(code.contains("import * as r from 'rkyv-js/decode';"));
1440 assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap/decode';"));
1441 assert!(code.contains("import { compileDecoder } from 'rkyv-js/jit/decode';"));
1443 assert!(code.contains("export const ArchivedEvent = compileDecoder(ArchivedEvent$);"));
1444 assert!(!code.contains("compileCodec"));
1445 }
1446
1447 #[test]
1448 fn set_jit_encode_direction_uses_compile_encoder() {
1449 let mut generator = CodeGenerator::new();
1450 generator.set_jit(true);
1451 generator.set_direction(Direction::Encode);
1452 generator.add_struct("Point", [("x", codec::f64())]);
1453 let code = generator.generate().unwrap();
1454 assert!(code.contains("import * as r from 'rkyv-js/encode';"));
1455 assert!(code.contains("import { compileEncoder } from 'rkyv-js/jit/encode';"));
1456 assert!(code.contains("export const ArchivedPoint = compileEncoder(ArchivedPoint$);"));
1457 }
1458}