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