use std::collections::{BTreeMap, BTreeSet};
use std::fs;
use std::path::Path;
use crate::casing::Casing;
use crate::error::{Diagnostic, DiagnosticKind, Error, SourceLocation};
use crate::expr::{CodecExpr, generate_import_block};
use crate::registry::{ExternalType, Registry, WithWrapper};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum OnUnknown {
#[default]
Error,
SkipContainingType,
}
#[derive(Debug, Clone)]
pub enum EnumVariant {
Unit(String),
Newtype(String, CodecExpr),
Tuple(String, Vec<CodecExpr>),
Struct(String, Vec<(String, CodecExpr)>),
}
impl EnumVariant {
pub fn name(&self) -> &str {
match self {
EnumVariant::Unit(name)
| EnumVariant::Newtype(name, _)
| EnumVariant::Tuple(name, _)
| EnumVariant::Struct(name, _) => name,
}
}
}
#[derive(Debug, Clone)]
pub(crate) enum TypeKind {
Struct(Vec<(String, CodecExpr)>),
Enum(Vec<EnumVariant>),
Alias(CodecExpr),
}
#[derive(Debug, Clone)]
struct FormatSpec {
endian: String,
pointer_width: u32,
aligned: bool,
}
impl FormatSpec {
fn is_default(&self) -> bool {
self.endian == "little" && self.pointer_width == 32 && self.aligned
}
fn options(&self) -> String {
let mut entries = Vec::new();
if self.endian != "little" {
entries.push(format!("endian: '{}'", self.endian));
}
if self.pointer_width != 32 {
entries.push(format!("pointerWidth: {}", self.pointer_width));
}
if !self.aligned {
entries.push("aligned: false".to_string());
}
entries.join(", ")
}
}
#[derive(Debug)]
pub struct CodeGenerator {
pub(crate) types: BTreeMap<String, TypeKind>,
pub(crate) failed: BTreeMap<String, Vec<Diagnostic>>,
pub(crate) add_diagnostics: Vec<Diagnostic>,
overrides: BTreeMap<String, String>,
header: Option<String>,
allow_typescript_syntax: bool,
pub(crate) on_unknown: OnUnknown,
pub(crate) marker_paths: BTreeSet<String>,
pub(crate) registry: Registry,
format: Option<FormatSpec>,
direction: Direction,
jit: bool,
field_casing: Casing,
variant_casing: Casing,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
pub enum Direction {
#[default]
Full,
Decode,
Encode,
}
impl Direction {
fn suffix(self) -> Option<&'static str> {
match self {
Direction::Full => None,
Direction::Decode => Some("decode"),
Direction::Encode => Some("encode"),
}
}
fn jit_entry(self) -> (&'static str, &'static str) {
match self {
Direction::Full => ("rkyv-js/jit", "compileCodec"),
Direction::Decode => ("rkyv-js/jit.decode", "compileDecoder"),
Direction::Encode => ("rkyv-js/jit.encode", "compileEncoder"),
}
}
fn split_specifier(self, spec: &str) -> Option<String> {
let suffix = self.suffix()?;
if spec == "rkyv-js" {
Some(format!("rkyv-js/{suffix}"))
} else if spec.starts_with("rkyv-js/lib/") {
Some(format!("{spec}.{suffix}"))
} else {
None
}
}
pub(crate) fn rewrite_import_block(self, block: &str) -> String {
if self == Direction::Full {
return block.to_string();
}
let mut out = String::with_capacity(block.len() + 64);
for line in block.lines() {
if let Some(spec_start) = line.rfind(" from '").map(|i| i + " from '".len())
&& let Some(len) = line[spec_start..].find('\'')
&& let Some(split) = self.split_specifier(&line[spec_start..spec_start + len])
{
out.push_str(&line[..spec_start]);
out.push_str(&split);
out.push_str(&line[spec_start + len..]);
out.push('\n');
continue;
}
out.push_str(line);
out.push('\n');
}
out
}
}
impl Default for CodeGenerator {
fn default() -> Self {
Self::new()
}
}
impl CodeGenerator {
pub fn new() -> Self {
Self {
types: BTreeMap::new(),
failed: BTreeMap::new(),
add_diagnostics: Vec::new(),
overrides: BTreeMap::new(),
header: None,
allow_typescript_syntax: true,
on_unknown: OnUnknown::Error,
marker_paths: BTreeSet::from(["rkyv::Archive".to_string()]),
registry: Registry::with_builtins(),
format: None,
direction: Direction::Full,
jit: false,
field_casing: Casing::Preserve,
variant_casing: Casing::Preserve,
}
}
pub fn set_header(&mut self, header: impl Into<String>) -> &mut Self {
self.header = Some(header.into());
self
}
pub fn set_direction(&mut self, direction: Direction) -> &mut Self {
self.direction = direction;
self
}
pub fn set_jit(&mut self, enabled: bool) -> &mut Self {
self.jit = enabled;
self
}
pub fn set_field_casing(&mut self, casing: Casing) -> &mut Self {
self.field_casing = casing;
self
}
pub fn set_variant_casing(&mut self, casing: Casing) -> &mut Self {
self.variant_casing = casing;
self
}
pub fn allow_typescript_syntax(&mut self, enabled: bool) -> &mut Self {
self.allow_typescript_syntax = enabled;
self
}
pub fn on_unknown_type(&mut self, mode: OnUnknown) -> &mut Self {
self.on_unknown = mode;
self
}
pub fn add_marker_path(&mut self, path: impl Into<String>) -> &mut Self {
self.marker_paths.insert(path.into());
self
}
pub fn register_external(
&mut self,
path: impl Into<String>,
external: ExternalType,
) -> &mut Self {
self.registry.register_type(path, external);
self
}
pub fn register_with(&mut self, path: impl Into<String>, wrapper: WithWrapper) -> &mut Self {
self.registry.register_wrapper(path, wrapper);
self
}
pub fn unregister_external(&mut self, path: &str) -> &mut Self {
self.registry.unregister_type(path);
self
}
pub fn add_struct(
&mut self,
name: impl Into<String>,
fields: impl IntoIterator<Item = (impl Into<String>, CodecExpr)>,
) -> &mut Self {
let fields = fields
.into_iter()
.map(|(field, expr)| (field.into(), expr))
.collect();
self.add_type(name.into(), TypeKind::Struct(fields), None);
self
}
pub fn add_enum(
&mut self,
name: impl Into<String>,
variants: impl IntoIterator<Item = EnumVariant>,
) -> &mut Self {
let variants = variants.into_iter().collect();
self.add_type(name.into(), TypeKind::Enum(variants), None);
self
}
pub fn add_alias(&mut self, name: impl Into<String>, target: CodecExpr) -> &mut Self {
self.add_type(name.into(), TypeKind::Alias(target), None);
self
}
pub(crate) fn add_type(
&mut self,
name: String,
kind: TypeKind,
location: Option<SourceLocation>,
) {
if self.is_known_type(&name) {
self.add_diagnostics.push(
Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
);
return;
}
self.types.insert(name, kind);
}
pub(crate) fn add_failed_type(
&mut self,
name: String,
diagnostics: Vec<Diagnostic>,
location: Option<SourceLocation>,
) {
if self.is_known_type(&name) {
self.add_diagnostics.push(
Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
);
return;
}
self.failed.insert(name, diagnostics);
}
fn is_known_type(&self, name: &str) -> bool {
self.types.contains_key(name) || self.failed.contains_key(name)
}
pub fn set_archived_name(
&mut self,
type_name: impl Into<String>,
archived_name: impl Into<String>,
) -> &mut Self {
self.overrides.insert(type_name.into(), archived_name.into());
self
}
pub fn archived_name_of(&self, type_name: &str) -> Option<String> {
if !self.is_known_type(type_name) {
return None;
}
Some(self.resolved_archived_name(type_name))
}
fn resolved_archived_name(&self, type_name: &str) -> String {
self.overrides
.get(type_name)
.cloned()
.unwrap_or_else(|| format!("Archived{type_name}"))
}
pub fn set_format(&mut self, endian: &str, pointer_width: u32, aligned: bool) -> &mut Self {
self.format = Some(FormatSpec {
endian: endian.to_string(),
pointer_width,
aligned,
});
self
}
fn nondefault_format(&self) -> Option<&FormatSpec> {
self.format.as_ref().filter(|spec| !spec.is_default())
}
fn exprs_with_context<'a>(
type_name: &str,
kind: &'a TypeKind,
) -> Vec<(String, &'a CodecExpr)> {
match kind {
TypeKind::Struct(fields) => fields
.iter()
.map(|(field, expr)| (format!("{type_name}.{field}"), expr))
.collect(),
TypeKind::Enum(variants) => {
let mut out = Vec::new();
for variant in variants {
match variant {
EnumVariant::Unit(_) => {}
EnumVariant::Newtype(vname, expr) => {
out.push((format!("{type_name}::{vname}"), expr));
}
EnumVariant::Tuple(vname, exprs) => {
for (i, expr) in exprs.iter().enumerate() {
out.push((format!("{type_name}::{vname}.{i}"), expr));
}
}
EnumVariant::Struct(vname, fields) => {
for (field, expr) in fields {
out.push((format!("{type_name}::{vname}.{field}"), expr));
}
}
}
}
out
}
TypeKind::Alias(expr) => vec![(type_name.to_string(), expr)],
}
}
fn casing_collisions(
context: &str,
names: impl IntoIterator<Item = String>,
casing: Casing,
) -> Vec<Diagnostic> {
if casing == Casing::Preserve {
return Vec::new();
}
let mut by_emitted: BTreeMap<String, Vec<String>> = BTreeMap::new();
for name in names {
by_emitted.entry(casing.apply(&name)).or_default().push(name);
}
by_emitted
.into_iter()
.filter(|(_, originals)| originals.len() > 1)
.map(|(emitted, originals)| {
Diagnostic::new(DiagnosticKind::NameCollision { emitted, originals })
.referenced_by(context.to_string())
})
.collect()
}
fn casing_diagnostics(&self, emitted: &BTreeMap<&String, &TypeKind>) -> Vec<Diagnostic> {
let mut diagnostics = Vec::new();
for (name, kind) in emitted {
match kind {
TypeKind::Struct(fields) => {
diagnostics.extend(Self::casing_collisions(
name,
fields.iter().map(|(field, _)| field.clone()),
self.field_casing,
));
}
TypeKind::Enum(variants) => {
diagnostics.extend(Self::casing_collisions(
name,
variants.iter().map(|variant| variant.name().to_string()),
self.variant_casing,
));
for variant in variants.iter() {
if let EnumVariant::Struct(vname, fields) = variant {
diagnostics.extend(Self::casing_collisions(
&format!("{name}::{vname}"),
fields.iter().map(|(field, _)| field.clone()),
self.field_casing,
));
}
}
}
TypeKind::Alias(_) => {}
}
}
diagnostics
}
pub fn generate(&self) -> Result<String, Error> {
let mut diagnostics: Vec<Diagnostic> = self.add_diagnostics.clone();
for target in self.overrides.keys() {
if !self.is_known_type(target) {
diagnostics.push(Diagnostic::new(DiagnosticKind::UnknownRenameTarget {
type_name: target.clone(),
}));
}
}
let mut skipped: BTreeSet<String> = BTreeSet::new();
match self.on_unknown {
OnUnknown::Error => {
for failure_diagnostics in self.failed.values() {
diagnostics.extend(failure_diagnostics.iter().cloned());
}
}
OnUnknown::SkipContainingType => {
for (name, failure_diagnostics) in &self.failed {
skipped.insert(name.clone());
for diagnostic in failure_diagnostics {
eprintln!(
"cargo:warning=rkyv-js-codegen: skipping `{name}`: {diagnostic}"
);
}
}
}
}
match self.on_unknown {
OnUnknown::Error => {
for (name, kind) in &self.types {
for (context, expr) in Self::exprs_with_context(name, kind) {
let mut refs = BTreeSet::new();
expr.collect_type_refs(&mut refs);
for reference in refs {
if !self.is_known_type(&reference) {
diagnostics.push(
Diagnostic::new(DiagnosticKind::UnresolvedTypeRef {
name: reference,
})
.referenced_by(context.clone()),
);
}
}
}
}
}
OnUnknown::SkipContainingType => {
loop {
let mut newly_skipped = Vec::new();
for (name, kind) in &self.types {
if skipped.contains(name) {
continue;
}
let broken = Self::exprs_with_context(name, kind).iter().any(
|(_, expr)| {
let mut refs = BTreeSet::new();
expr.collect_type_refs(&mut refs);
refs.iter().any(|reference| {
skipped.contains(reference)
|| !self.types.contains_key(reference)
})
},
);
if broken {
newly_skipped.push(name.clone());
}
}
if newly_skipped.is_empty() {
break;
}
for name in newly_skipped {
eprintln!(
"cargo:warning=rkyv-js-codegen: skipping `{name}`: it references \
a type that was omitted or never added"
);
skipped.insert(name);
}
}
}
}
let emitted: BTreeMap<&String, &TypeKind> = self
.types
.iter()
.filter(|(name, _)| !skipped.contains(*name))
.collect();
diagnostics.extend(self.casing_diagnostics(&emitted));
let (jit_module, jit_fn) = self.direction.jit_entry();
let jit_import = CodecExpr::import_from(jit_module, jit_fn);
let mut all_exprs: Vec<&CodecExpr> = emitted
.iter()
.flat_map(|(name, kind)| Self::exprs_with_context(name, kind))
.map(|(_, expr)| expr)
.collect();
if self.jit && !emitted.is_empty() {
all_exprs.push(&jit_import);
}
let import_block = match generate_import_block(all_exprs.iter().copied()) {
Ok(block) => self.direction.rewrite_import_block(&block),
Err(conflicts) => {
diagnostics.extend(conflicts.into_iter().map(Diagnostic::new));
String::new()
}
};
if !diagnostics.is_empty() {
return Err(Error::Codegen(diagnostics));
}
let order = Self::topological_sort(&emitted);
let archived_names: BTreeMap<String, String> = emitted
.keys()
.map(|name| ((*name).clone(), self.resolved_archived_name(name)))
.collect();
let codec_names: BTreeMap<String, String> = if self.jit {
archived_names
.iter()
.map(|(name, archived)| (name.clone(), format!("{archived}$")))
.collect()
} else {
archived_names.clone()
};
let mut blocks: Vec<String> = Vec::new();
let header = self
.header
.as_deref()
.unwrap_or("Auto-generated by rkyv-js-codegen\nDO NOT EDIT MANUALLY");
let mut header_block = String::from("/**\n");
for line in header.lines() {
if line.is_empty() {
header_block.push_str(" *\n");
} else {
header_block.push_str(" * ");
header_block.push_str(line);
header_block.push('\n');
}
}
header_block.push_str(" */");
blocks.push(header_block);
blocks.push(import_block.trim_end().to_string());
if let Some(spec) = self.nondefault_format() {
blocks.push(format!("const FORMAT = r.format({{ {} }});", spec.options()));
}
for name in &order {
let kind = emitted.get(name).expect("ordered names come from emitted");
blocks.push(self.emit_type(name, kind, &archived_names, &codec_names));
}
Ok(blocks.join("\n\n") + "\n")
}
fn topological_sort(emitted: &BTreeMap<&String, &TypeKind>) -> Vec<String> {
let mut deps: BTreeMap<&str, BTreeSet<String>> = BTreeMap::new();
for (name, kind) in emitted {
let mut refs = BTreeSet::new();
for (_, expr) in Self::exprs_with_context(name, kind) {
expr.collect_type_refs(&mut refs);
}
refs.retain(|reference| {
emitted.contains_key(reference) && reference != name.as_str()
});
deps.insert(name.as_str(), refs);
}
let mut dependents: BTreeMap<&str, Vec<&str>> = BTreeMap::new();
let mut in_degree: BTreeMap<&str, usize> = BTreeMap::new();
for (name, type_deps) in &deps {
in_degree.insert(name, type_deps.len());
for dep in type_deps {
dependents.entry(dep.as_str()).or_default().push(name);
}
}
let mut ready: BTreeSet<&str> = in_degree
.iter()
.filter(|(_, degree)| **degree == 0)
.map(|(name, _)| *name)
.collect();
let mut order: Vec<String> = Vec::new();
let mut done: BTreeSet<&str> = BTreeSet::new();
while let Some(name) = ready.pop_first() {
order.push(name.to_string());
done.insert(name);
if let Some(children) = dependents.get(name) {
for child in children {
let degree = in_degree.get_mut(child).unwrap();
*degree -= 1;
if *degree == 0 {
ready.insert(child);
}
}
}
}
for name in deps.keys() {
if !done.contains(name) {
order.push((*name).to_string());
}
}
order
}
fn emit_type(
&self,
name: &str,
kind: &TypeKind,
archived_names: &BTreeMap<String, String>,
codec_names: &BTreeMap<String, String>,
) -> String {
let archived = archived_names
.get(name)
.expect("emitted types have archived names")
.clone();
let render = |expr: &CodecExpr| -> String {
expr.render(codec_names)
.expect("type references are validated before emission")
};
let codec_expr = match kind {
TypeKind::Struct(fields) => {
if fields.is_empty() {
"r.struct({})".to_string()
} else {
let mut body = String::from("r.struct({\n");
for (field, expr) in fields {
body.push_str(&format!(
" {}: {},\n",
self.field_casing.apply(field),
render(expr)
));
}
body.push_str("})");
body
}
}
TypeKind::Enum(variants) => {
if variants.is_empty() {
"r.taggedEnum({})".to_string()
} else {
let mut body = String::from("r.taggedEnum({\n");
for variant in variants {
let value = match variant {
EnumVariant::Unit(_) => "null".to_string(),
EnumVariant::Newtype(_, expr) => render(expr),
EnumVariant::Tuple(_, exprs) => {
render(&CodecExpr::array(exprs.iter().cloned()))
}
EnumVariant::Struct(_, fields) => {
let record = CodecExpr::object(fields.iter().map(
|(field, expr)| {
(self.field_casing.apply(field), expr.clone())
},
));
render(&record)
}
};
body.push_str(&format!(
" {}: {},\n",
self.variant_casing.apply(variant.name()),
value
));
}
body.push_str("})");
body
}
}
TypeKind::Alias(expr) => render(expr),
};
let codec_expr = match self.nondefault_format() {
Some(_) => format!("r.withFormat({codec_expr}, FORMAT)"),
None => codec_expr,
};
let mut block = if self.jit {
let jit_fn = self.direction.jit_entry().1;
format!(
"const {archived}$ = {codec_expr};\n\n\
export const {archived} = {jit_fn}({archived}$);"
)
} else {
format!("export const {archived} = {codec_expr};")
};
if self.allow_typescript_syntax {
block.push_str(&format!(
"\n\nexport type {name} = r.Infer<typeof {archived}>;"
));
}
block
}
pub fn write_to_file(&self, path: impl AsRef<Path>) -> Result<(), Error> {
let code = self.generate()?;
fs::write(path, code)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::expr::codec;
fn diagnostics(error: Error) -> Vec<Diagnostic> {
match error {
Error::Codegen(diagnostics) => diagnostics,
other => panic!("expected Error::Codegen, got {other:?}"),
}
}
#[test]
fn struct_emission_snapshot() {
let mut generator = CodeGenerator::new();
generator.add_struct("Point", [("x", codec::f64()), ("y", codec::f64())]);
let code = generator.generate().unwrap();
assert_eq!(
code,
"/**\n\
\x20* Auto-generated by rkyv-js-codegen\n\
\x20* DO NOT EDIT MANUALLY\n\
\x20*/\n\
\n\
import * as r from 'rkyv-js';\n\
\n\
export const ArchivedPoint = r.struct({\n\
\x20 x: r.f64,\n\
\x20 y: r.f64,\n\
});\n\
\n\
export type Point = r.Infer<typeof ArchivedPoint>;\n"
);
}
#[test]
fn enum_emission_snapshot() {
let mut generator = CodeGenerator::new();
generator.add_enum(
"MixedAlign",
[
EnumVariant::Struct(
"V".to_string(),
vec![("a".to_string(), codec::u8()), ("b".to_string(), codec::u32())],
),
EnumVariant::Newtype("X".to_string(), codec::u64()),
EnumVariant::Tuple("Color".to_string(), vec![codec::u8(), codec::u8()]),
EnumVariant::Unit("Y".to_string()),
],
);
let code = generator.generate().unwrap();
assert!(code.contains(
"export const ArchivedMixedAlign = r.taggedEnum({\n\
\x20 V: { a: r.u8, b: r.u32 },\n\
\x20 X: r.u64,\n\
\x20 Color: [r.u8, r.u8],\n\
\x20 Y: null,\n\
});"
));
assert!(code.contains("export type MixedAlign = r.Infer<typeof ArchivedMixedAlign>;"));
}
#[test]
fn alias_emission_snapshot() {
let mut generator = CodeGenerator::new();
generator.add_alias("UserId", codec::u32());
let code = generator.generate().unwrap();
assert!(code.contains("export const ArchivedUserId = r.u32;"));
assert!(code.contains("export type UserId = r.Infer<typeof ArchivedUserId>;"));
}
#[test]
fn imports_are_collected_and_deduped() {
let mut generator = CodeGenerator::new();
generator.add_struct(
"A",
[
(
"m",
CodecExpr::call(
CodecExpr::import_from("rkyv-js/lib/hashmap", "hashMap"),
[codec::string(), codec::u32()],
),
),
(
"s",
CodecExpr::call(
CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
[codec::string()],
),
),
],
);
generator.add_struct(
"B",
[(
"s2",
CodecExpr::call(
CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
[codec::u32()],
),
)],
);
let code = generator.generate().unwrap();
assert!(code.contains("import { hashMap, hashSet } from 'rkyv-js/lib/hashmap';"));
assert_eq!(code.matches("hashSet }").count(), 1);
}
#[test]
fn import_conflict_is_reported() {
let mut generator = CodeGenerator::new();
generator.add_struct("A", [("x", CodecExpr::import_from("pkg-a", "codec"))]);
generator.add_struct("B", [("y", CodecExpr::import_from("pkg-b", "codec"))]);
let errors = diagnostics(generator.generate().unwrap_err());
assert!(errors.iter().any(|diagnostic| matches!(
&diagnostic.kind,
DiagnosticKind::ImportConflict { export, .. } if export == "codec"
)));
}
#[test]
fn topo_sort_handles_forward_references() {
let mut generator = CodeGenerator::new();
generator.add_struct("AOuter", [("inner", codec::named("Inner"))]);
generator.add_struct("Inner", [("value", codec::u32())]);
let code = generator.generate().unwrap();
let inner_pos = code.find("export const ArchivedInner").unwrap();
let outer_pos = code.find("export const ArchivedAOuter").unwrap();
assert!(inner_pos < outer_pos, "dependency must be emitted first");
assert!(code.contains("inner: ArchivedInner,"));
}
#[test]
fn unresolved_type_ref_reports_referrer() {
let mut generator = CodeGenerator::new();
generator.add_struct("Outer", [("inner", codec::named("Missing"))]);
let errors = diagnostics(generator.generate().unwrap_err());
assert_eq!(errors.len(), 1);
assert!(matches!(
&errors[0].kind,
DiagnosticKind::UnresolvedTypeRef { name } if name == "Missing"
));
assert_eq!(errors[0].referenced_by.as_deref(), Some("Outer.inner"));
}
#[test]
fn duplicate_type_is_reported_at_generate() {
let mut generator = CodeGenerator::new();
generator.add_struct("Point", [("x", codec::f64())]);
generator.add_struct("Point", [("y", codec::f64())]);
let errors = diagnostics(generator.generate().unwrap_err());
assert!(errors.iter().any(|diagnostic| matches!(
&diagnostic.kind,
DiagnosticKind::DuplicateType { name } if name == "Point"
)));
}
#[test]
fn set_archived_name_is_order_independent() {
let mut generator = CodeGenerator::new();
generator.set_archived_name("Foo", "MyFoo");
generator.add_struct("Foo", [("x", codec::u32())]);
let code = generator.generate().unwrap();
assert!(code.contains("export const MyFoo = r.struct({"));
assert!(code.contains("export type Foo = r.Infer<typeof MyFoo>;"));
assert!(!code.contains("ArchivedFoo"));
let mut generator = CodeGenerator::new();
generator.add_struct("Foo", [("x", codec::u32())]);
generator.set_archived_name("Foo", "MyFoo");
let code = generator.generate().unwrap();
assert!(code.contains("export const MyFoo = r.struct({"));
}
#[test]
fn archived_rename_applies_to_cross_references() {
let mut generator = CodeGenerator::new();
generator.set_archived_name("Inner", "CustomInner");
generator.add_struct("Inner", [("value", codec::u32())]);
generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
let code = generator.generate().unwrap();
assert!(code.contains("export const CustomInner = r.struct({"));
assert!(code.contains("inner: CustomInner,"));
assert!(!code.contains("ArchivedInner"));
}
#[test]
fn unknown_rename_target_is_a_diagnostic() {
let mut generator = CodeGenerator::new();
generator.add_struct("Foo", [("x", codec::u32())]);
generator.set_archived_name("Nope", "MyNope");
let errors = diagnostics(generator.generate().unwrap_err());
assert!(errors.iter().any(|diagnostic| matches!(
&diagnostic.kind,
DiagnosticKind::UnknownRenameTarget { type_name } if type_name == "Nope"
)));
}
#[test]
fn archived_name_of_accessor() {
let mut generator = CodeGenerator::new();
generator.add_struct("Foo", [("x", codec::u32())]);
assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("ArchivedFoo"));
generator.set_archived_name("Foo", "MyFoo");
assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("MyFoo"));
assert_eq!(generator.archived_name_of("Bar"), None);
}
#[test]
fn field_casing_defaults_to_preserve() {
let mut generator = CodeGenerator::new();
generator.add_struct("Event", [("created_at", codec::u64())]);
let code = generator.generate().unwrap();
assert!(code.contains("created_at: r.u64,"));
}
#[test]
fn field_casing_camel_rewrites_struct_fields() {
let mut generator = CodeGenerator::new();
generator.set_field_casing(Casing::Camel);
generator.add_struct(
"Event",
[
("created_at", codec::u64()),
("HTTP_status", codec::u16()),
("id", codec::u32()),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("createdAt: r.u64,"));
assert!(code.contains("httpStatus: r.u16,"));
assert!(code.contains("id: r.u32,"));
let created = code.find("createdAt").unwrap();
let status = code.find("httpStatus").unwrap();
assert!(created < status);
}
#[test]
fn field_casing_applies_to_enum_struct_variants() {
let mut generator = CodeGenerator::new();
generator.set_field_casing(Casing::Camel);
generator.add_enum(
"Message",
[
EnumVariant::Struct(
"Text".to_string(),
vec![
("sent_at".to_string(), codec::u64()),
("body_text".to_string(), codec::string()),
],
),
EnumVariant::Unit("Ping".to_string()),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("Text: { sentAt: r.u64, bodyText: r.string },"));
assert!(code.contains("Ping: null,"));
}
#[test]
fn variant_casing_is_independent_of_field_casing() {
let mut generator = CodeGenerator::new();
generator
.set_field_casing(Casing::Camel)
.set_variant_casing(Casing::Snake);
generator.add_enum(
"Message",
[
EnumVariant::Struct(
"PlainText".to_string(),
vec![("sent_at".to_string(), codec::u64())],
),
EnumVariant::Newtype("BinaryBlob".to_string(), codec::string()),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("plain_text: { sentAt: r.u64 },"));
assert!(code.contains("binary_blob: r.string,"));
}
#[test]
fn casing_leaves_type_and_export_names_alone() {
let mut generator = CodeGenerator::new();
generator.set_field_casing(Casing::Camel);
generator.add_struct("HttpEvent", [("created_at", codec::u64())]);
let code = generator.generate().unwrap();
assert!(code.contains("export const ArchivedHttpEvent = r.struct({"));
assert!(code.contains("export type HttpEvent = r.Infer<typeof ArchivedHttpEvent>;"));
}
#[test]
fn casing_collision_is_reported() {
let mut generator = CodeGenerator::new();
generator.set_field_casing(Casing::Camel);
generator.add_struct(
"Event",
[("foo_bar", codec::u32()), ("fooBar", codec::u32())],
);
let errors = diagnostics(generator.generate().unwrap_err());
assert_eq!(errors.len(), 1);
assert!(matches!(
&errors[0].kind,
DiagnosticKind::NameCollision { emitted, originals }
if emitted == "fooBar" && originals.len() == 2
));
assert_eq!(errors[0].referenced_by.as_deref(), Some("Event"));
}
#[test]
fn casing_collision_in_a_struct_variant_names_the_variant() {
let mut generator = CodeGenerator::new();
generator.set_field_casing(Casing::Camel);
generator.add_enum(
"Message",
[EnumVariant::Struct(
"Text".to_string(),
vec![
("sent_at".to_string(), codec::u64()),
("sentAt".to_string(), codec::u64()),
],
)],
);
let errors = diagnostics(generator.generate().unwrap_err());
assert_eq!(errors.len(), 1);
assert_eq!(errors[0].referenced_by.as_deref(), Some("Message::Text"));
}
#[test]
fn variant_casing_collision_is_reported() {
let mut generator = CodeGenerator::new();
generator.set_variant_casing(Casing::Snake);
generator.add_enum(
"Message",
[
EnumVariant::Unit("PlainText".to_string()),
EnumVariant::Unit("plain_text".to_string()),
],
);
let errors = diagnostics(generator.generate().unwrap_err());
assert!(errors.iter().any(|diagnostic| matches!(
&diagnostic.kind,
DiagnosticKind::NameCollision { emitted, .. } if emitted == "plain_text"
)));
}
#[test]
fn preserve_never_reports_a_collision() {
let mut generator = CodeGenerator::new();
generator.add_struct(
"Event",
[("foo_bar", codec::u32()), ("fooBar", codec::u32())],
);
let code = generator.generate().unwrap();
assert!(code.contains("foo_bar: r.u32,"));
assert!(code.contains("fooBar: r.u32,"));
}
#[test]
fn js_mode_omits_type_lines() {
let mut generator = CodeGenerator::new();
generator.allow_typescript_syntax(false);
generator.add_struct("Point", [("x", codec::f64())]);
generator.add_alias("UserId", codec::u32());
let code = generator.generate().unwrap();
assert!(code.contains("export const ArchivedPoint = r.struct({"));
assert!(code.contains("export const ArchivedUserId = r.u32;"));
assert!(!code.contains("export type"));
assert!(!code.contains("r.Infer"));
}
#[test]
fn set_format_default_is_a_no_op() {
let mut generator = CodeGenerator::new();
generator.set_format("little", 32, true);
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(!code.contains("FORMAT"));
assert!(!code.contains("withFormat"));
}
#[test]
fn set_format_nondefault_wraps_exports() {
let mut generator = CodeGenerator::new();
generator.set_format("big", 64, false);
generator.add_struct("Point", [("x", codec::f64())]);
generator.add_alias("UserId", codec::u32());
let code = generator.generate().unwrap();
assert!(code.contains(
"const FORMAT = r.format({ endian: 'big', pointerWidth: 64, aligned: false });"
));
assert!(code.contains("export const ArchivedPoint = r.withFormat(r.struct({\n"));
assert!(code.contains("}), FORMAT);"));
assert!(code.contains("export const ArchivedUserId = r.withFormat(r.u32, FORMAT);"));
}
#[test]
fn set_format_emits_only_nondefault_keys() {
let mut generator = CodeGenerator::new();
generator.set_format("little", 16, true);
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(code.contains("const FORMAT = r.format({ pointerWidth: 16 });"));
}
#[test]
fn custom_header_replaces_default() {
let mut generator = CodeGenerator::new();
generator.set_header("Custom header\nsecond line");
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(code.starts_with("/**\n * Custom header\n * second line\n */\n"));
assert!(!code.contains("Auto-generated by rkyv-js-codegen"));
}
#[test]
fn set_direction_full_is_a_no_op() {
let mut generator = CodeGenerator::new();
generator.set_direction(Direction::Full);
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(code.contains("import * as r from 'rkyv-js';"));
}
#[test]
fn set_direction_rewrites_rkyv_specifiers_only() {
let mut generator = CodeGenerator::new();
generator.set_direction(Direction::Decode);
generator.add_struct(
"Event",
[
("id", codec::u32()),
(
"tags",
CodecExpr::call(
CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
[codec::string()],
),
),
(
"custom",
CodecExpr::import_from("./my-codec.ts", "MyCodec"),
),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("import * as r from 'rkyv-js/decode';"));
assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap.decode';"));
assert!(code.contains("import { MyCodec } from './my-codec.ts';"));
assert!(code.contains("export const ArchivedEvent = r.struct({"));
assert!(code.contains("export type Event = r.Infer<typeof ArchivedEvent>;"));
}
#[test]
fn set_direction_encode_uses_encode_suffix() {
let mut generator = CodeGenerator::new();
generator.set_direction(Direction::Encode);
generator.add_struct(
"Point",
[
("x", codec::f64()),
("id", CodecExpr::import_from("rkyv-js/lib/uuid", "uuid")),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("import * as r from 'rkyv-js/encode';"));
assert!(code.contains("import { uuid } from 'rkyv-js/lib/uuid.encode';"));
}
#[test]
fn split_specifiers_mirror_the_runtime_module_names() {
for (direction, root, lib) in [
(Direction::Decode, "rkyv-js/decode", "rkyv-js/lib/bytes.decode"),
(Direction::Encode, "rkyv-js/encode", "rkyv-js/lib/bytes.encode"),
] {
let mut generator = CodeGenerator::new();
generator.set_direction(direction);
generator.add_struct(
"Blob",
[
("len", codec::u32()),
("data", CodecExpr::import_from("rkyv-js/lib/bytes", "bytes")),
],
);
let code = generator.generate().unwrap();
assert!(code.contains(&format!("import * as r from '{root}';")));
assert!(code.contains(&format!("import {{ bytes }} from '{lib}';")));
}
}
#[test]
fn set_jit_wraps_exports() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.add_struct("Point", [("x", codec::f64())]);
generator.add_alias("UserId", codec::u32());
let code = generator.generate().unwrap();
assert!(code.contains("import { compileCodec } from 'rkyv-js/jit';"));
assert!(code.contains("const ArchivedPoint$ = r.struct({\n"));
assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
assert!(code.contains("const ArchivedUserId$ = r.u32;"));
assert!(code.contains("export const ArchivedUserId = compileCodec(ArchivedUserId$);"));
assert!(code.contains("export type Point = r.Infer<typeof ArchivedPoint>;"));
}
#[test]
fn set_jit_references_resolve_to_raw_codecs() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.add_struct("Inner", [("value", codec::u32())]);
generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
let code = generator.generate().unwrap();
assert!(code.contains("inner: ArchivedInner$,"));
assert!(code.contains("export const ArchivedInner = compileCodec(ArchivedInner$);"));
assert!(code.contains("export const ArchivedOuter = compileCodec(ArchivedOuter$);"));
}
#[test]
fn set_jit_composes_with_format() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.set_format("big", 64, true);
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(code.contains("const FORMAT = r.format({ endian: 'big', pointerWidth: 64 });"));
assert!(code.contains("const ArchivedPoint$ = r.withFormat(r.struct({\n"));
assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
}
#[test]
fn set_jit_respects_archived_renames() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.set_archived_name("Inner", "CustomInner");
generator.add_struct("Inner", [("value", codec::u32())]);
generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
let code = generator.generate().unwrap();
assert!(code.contains("inner: CustomInner$,"));
assert!(code.contains("export const CustomInner = compileCodec(CustomInner$);"));
}
#[test]
fn set_jit_decode_direction_uses_compile_decoder() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.set_direction(Direction::Decode);
generator.add_struct(
"Event",
[
("id", codec::u32()),
(
"tags",
CodecExpr::call(
CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
[codec::string()],
),
),
],
);
let code = generator.generate().unwrap();
assert!(code.contains("import * as r from 'rkyv-js/decode';"));
assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap.decode';"));
assert!(code.contains("import { compileDecoder } from 'rkyv-js/jit.decode';"));
assert!(!code.contains("jit.decode.decode"));
assert!(code.contains("export const ArchivedEvent = compileDecoder(ArchivedEvent$);"));
assert!(!code.contains("compileCodec"));
}
#[test]
fn set_jit_encode_direction_uses_compile_encoder() {
let mut generator = CodeGenerator::new();
generator.set_jit(true);
generator.set_direction(Direction::Encode);
generator.add_struct("Point", [("x", codec::f64())]);
let code = generator.generate().unwrap();
assert!(code.contains("import * as r from 'rkyv-js/encode';"));
assert!(code.contains("import { compileEncoder } from 'rkyv-js/jit.encode';"));
assert!(code.contains("export const ArchivedPoint = compileEncoder(ArchivedPoint$);"));
}
}