Skip to main content

rkyv_js_codegen/
generator.rs

1//! The TypeScript binding generator.
2
3use 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/// How to handle a field whose type cannot be mapped to a codec.
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
14pub enum OnUnknown {
15    /// Aggregate a diagnostic and fail [`generate`](CodeGenerator::generate).
16    #[default]
17    Error,
18    /// Emit a `cargo:warning` and omit the containing type — and,
19    /// transitively, every type referencing it — from the output.
20    SkipContainingType,
21}
22
23/// An enum variant for [`CodeGenerator::add_enum`].
24#[derive(Debug, Clone)]
25pub enum EnumVariant {
26    /// A unit variant: `Name` - emitted as `Name: null`.
27    Unit(String),
28    /// A newtype (1-tuple) variant: `Name(T)` - emitted as a bare codec.
29    Newtype(String, CodecExpr),
30    /// An n-tuple variant (n >= 2): `Name(T0, T1)` - emitted as an array of codecs (`[t0, t1]`), decoded as an array value.
31    /// The fields stay flattened in the enum layout (this is NOT a nested `r.tuple` block).
32    Tuple(String, Vec<CodecExpr>),
33    /// A struct variant: `Name { a: T }` — emitted as a record of codecs.
34    Struct(String, Vec<(String, CodecExpr)>),
35}
36
37impl EnumVariant {
38    /// The variant name.
39    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/// The kind-specific payload of a generated type.
50#[derive(Debug, Clone)]
51pub(crate) enum TypeKind {
52    Struct(Vec<(String, CodecExpr)>),
53    Enum(Vec<EnumVariant>),
54    Alias(CodecExpr),
55}
56
57/// The non-default wire format configured via [`set_format`](CodeGenerator::set_format).
58#[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    /// The non-default keys as `r.format(...)` options.
71    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/// Collects type definitions — from Rust sources or programmatically — and
87/// generates TypeScript codec bindings for the `rkyv-js` runtime.
88///
89/// # Example
90///
91/// ```
92/// use rkyv_js_codegen::{CodeGenerator, codec};
93///
94/// let mut generator = CodeGenerator::new();
95/// generator.add_struct("Point", [("x", codec::f64()), ("y", codec::f64())]);
96/// let code = generator.generate().unwrap();
97/// assert!(code.contains("export const ArchivedPoint = r.struct({"));
98/// assert!(code.contains("export type Point = r.Infer<typeof ArchivedPoint>;"));
99/// ```
100#[derive(Debug)]
101pub struct CodeGenerator {
102    /// Successfully added types, keyed by Rust type name.
103    pub(crate) types: BTreeMap<String, TypeKind>,
104    /// Types whose extraction produced diagnostics, keyed by Rust type name.
105    pub(crate) failed: BTreeMap<String, Vec<Diagnostic>>,
106    /// Diagnostics recorded at add time (duplicate type names).
107    pub(crate) add_diagnostics: Vec<Diagnostic>,
108    /// `set_archived_name` overrides, applied at generate time.
109    overrides: BTreeMap<String, String>,
110    header: Option<String>,
111    allow_typescript_syntax: bool,
112    pub(crate) on_unknown: OnUnknown,
113    /// Derive paths that mark a type for extraction.
114    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/// Which half of the codec surface the generated bindings target.
124///
125/// The emitted factory calls and type exports are identical in all three modes.
126/// Only the `rkyv-js` import specifiers change, so a decode-only bundle never pulls the writer/hasher machinery (and vice versa).
127#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
128pub enum Direction {
129    /// Full codecs (`rkyv-js`): encode + decode + access.
130    #[default]
131    Full,
132    /// Decoder-only bindings (`rkyv-js/decode`, `rkyv-js/lib/*.decode`).
133    Decode,
134    /// Encoder-only bindings (`rkyv-js/encode`, `rkyv-js/lib/*.encode`).
135    Encode,
136}
137
138impl Direction {
139    /// The module basename this direction's entry points carry.
140    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    /// The JIT entry point and compile function for this direction.
149    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    /// This direction's counterpart of an `rkyv-js` specifier, or `None` when
158    /// the runtime does not split that module.
159    ///
160    /// Every split module sits next to the one it splits, so the specifier
161    /// mirrors the file name: `rkyv-js/lib/hashmap` pairs with
162    /// `rkyv-js/lib/hashmap.decode`. The package root is the one exception —
163    /// it resolves to `index`, whose counterpart is the separate `decode`
164    /// module, hence `rkyv-js/decode`.
165    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    /// Rewrite an emitted import block's `rkyv-js` specifiers for this direction.
177    /// Non-`rkyv-js` specifiers (user `register_external` modules) are left untouched.
178    /// Hand-written codecs must provide their own direction-appropriate exports.
179    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    /// Create a generator with the built-in type and wrapper registrations.
210    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    /// Replace the header comment of the generated file.
230    pub fn set_header(&mut self, header: impl Into<String>) -> &mut Self {
231        self.header = Some(header.into());
232        self
233    }
234
235    /// Emit unidirectional bindings: [`Direction::Decode`] rewrites every `rkyv-js` import specifier
236    /// to its decode counterpart (`rkyv-js` becomes `rkyv-js/decode`, `rkyv-js/lib/X` becomes
237    /// `rkyv-js/lib/X.decode`), [`Direction::Encode`] symmetrically.
238    ///
239    /// Factory names and type exports are unchanged;
240    /// imports of user modules registered via `register_external` are not rewritten.
241    pub fn set_direction(&mut self, direction: Direction) -> &mut Self {
242        self.direction = direction;
243        self
244    }
245
246    /// Wrap every exported codec in the direction-matched JIT compile function: `compileCodec` from `rkyv-js/jit` for [`Direction::Full`],
247    /// `compileDecoder` from `rkyv-js/jit.decode` resp. `compileEncoder` from `rkyv-js/jit.encode` for unidirectional bindings.
248    ///
249    /// Each type is emitted as a non-exported interpreter codec (`const {Name}$ = ...`)
250    /// plus a compiled export (`export const {Name} = compileCodec({Name}$);`),
251    /// and cross-references between generated types resolve to the `$` codecs:
252    /// a compiled codec is opaque to the JIT, so compiling each export over the
253    /// raw graph is what lets nested types inline instead of degrading to
254    /// per-element dispatch calls. The compiled exports stay drop-in
255    /// (`encode`/`decode`/`access`/... and `r.Infer` are unchanged),
256    /// and fall back to the interpreter codec where `new Function` is blocked (CSP).
257    ///
258    /// Every export compiles eagerly at module load.
259    ///
260    /// Defaults to `false`.
261    pub fn set_jit(&mut self, enabled: bool) -> &mut Self {
262        self.jit = enabled;
263        self
264    }
265
266    /// Rewrite the casing of emitted struct field names — including the fields
267    /// of enum struct variants — so the decoded objects read as idiomatic
268    /// JavaScript: `Casing::Camel` turns Rust's `created_at` into `createdAt`.
269    ///
270    /// rkyv lays a struct out positionally, so the keys of the emitted
271    /// `r.struct({ ... })` are labels only. Renaming them changes the shape of
272    /// the decoded object and the inferred `r.Infer` type, and does not move a
273    /// single wire byte: bindings generated with and without this option stay
274    /// interchangeable on the same buffer.
275    ///
276    /// Names that collide after conversion (`foo_bar` and `fooBar` both
277    /// becoming `fooBar`) are reported as [`DiagnosticKind::NameCollision`]
278    /// rather than emitted as a duplicate object key.
279    ///
280    /// Defaults to [`Casing::Preserve`].
281    ///
282    /// ```
283    /// use rkyv_js_codegen::{Casing, CodeGenerator, codec};
284    ///
285    /// let mut generator = CodeGenerator::new();
286    /// generator.set_field_casing(Casing::Camel);
287    /// generator.add_struct("Event", [("created_at", codec::u64())]);
288    /// assert!(generator.generate()?.contains("createdAt: r.u64,"));
289    /// # Ok::<(), rkyv_js_codegen::Error>(())
290    /// ```
291    pub fn set_field_casing(&mut self, casing: Casing) -> &mut Self {
292        self.field_casing = casing;
293        self
294    }
295
296    /// Rewrite the casing of emitted enum variant names — the keys of
297    /// `r.taggedEnum({ ... })`, which surface as the `tag` of every decoded
298    /// value.
299    ///
300    /// Rust variants are already `PascalCase`, which is the conventional
301    /// spelling for a discriminated-union tag in TypeScript, so this is
302    /// separate from [`set_field_casing`](Self::set_field_casing) and defaults
303    /// to [`Casing::Preserve`].
304    ///
305    /// The discriminant on the wire is the variant's index, not its name, so
306    /// this is a relabelling just like `set_field_casing`.
307    pub fn set_variant_casing(&mut self, casing: Casing) -> &mut Self {
308        self.variant_casing = casing;
309        self
310    }
311
312    /// When `false`, `export type ... = r.Infer<...>` lines are dropped so the output is valid plain JavaScript.
313    ///
314    /// Defaults to `true`.
315    pub fn allow_typescript_syntax(&mut self, enabled: bool) -> &mut Self {
316        self.allow_typescript_syntax = enabled;
317        self
318    }
319
320    /// Configure how unmappable field types are handled.
321    ///
322    /// Defaults to [`OnUnknown::Error`].
323    pub fn on_unknown_type(&mut self, mode: OnUnknown) -> &mut Self {
324        self.on_unknown = mode;
325        self
326    }
327
328    /// Register an additional derive path that marks types for extraction,
329    /// alongside the default `rkyv::Archive`.
330    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    /// Register (or replace) an external type mapping for a fully-qualified Rust path.
336    ///
337    /// ```
338    /// use rkyv_js_codegen::{CodeGenerator, CodecExpr, ExternalType};
339    ///
340    /// let mut generator = CodeGenerator::new();
341    /// generator.register_external(
342    ///     "my_crate::MyVec",
343    ///     ExternalType::generic1(|t| {
344    ///         CodecExpr::call(CodecExpr::import_from("my-pkg/codecs", "myVec"), [t])
345    ///     }),
346    /// );
347    /// ```
348    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    /// Register (or replace) a `#[rkyv(with = ...)]` wrapper handler.
358    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    /// Remove an external type mapping (e.g. to disable a builtin).
364    pub fn unregister_external(&mut self, path: &str) -> &mut Self {
365        self.registry.unregister_type(path);
366        self
367    }
368
369    /// Add a struct definition.
370    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    /// Add an enum definition.
384    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    /// Add a type alias: `export const Archived{name} = <expr>;`.
395    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    /// Record a type entry, diagnosing duplicate names.
401    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    /// Record a type whose extraction produced diagnostics.
417    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    /// Override the archived (exported) name of a type, corresponding to `#[rkyv(archived = Name)]`.
437    ///
438    /// Order-independent: the target type may be added before or after this call.
439    /// A target that never materializes is reported as [`DiagnosticKind::UnknownRenameTarget`] at generate time.
440    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    /// The archived (exported) name a type will be emitted under, or `None`
450    /// if no type with that name has been added.
451    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    /// Configure the rkyv wire format of the generated bindings.
466    ///
467    /// When the format differs from the default (`little`/32/aligned),
468    /// the output declares `const FORMAT = r.format({ ... })` with the non-default keys
469    /// and wraps every exported codec in `r.withFormat(<expr>, FORMAT)`.
470    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    /// The active non-default format, if any.
480    fn nondefault_format(&self) -> Option<&FormatSpec> {
481        self.format.as_ref().filter(|spec| !spec.is_default())
482    }
483
484    /// Every codec expression of a type, labelled with its `Type.field`
485    /// provenance for diagnostics.
486    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    /// Collisions among `names` after `casing` conversion, as diagnostics
522    /// tagged with `context`.
523    ///
524    /// Emitted names key a JavaScript object literal, so a collision would not
525    /// fail loudly — it would drop a field and shift every offset after it.
526    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    /// Every casing collision across the types that will be emitted.
549    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    /// Generate the TypeScript bindings.
583    ///
584    /// Validation runs first; every problem is aggregated into a single [`Error::Codegen`].
585    pub fn generate(&self) -> Result<String, Error> {
586        let mut diagnostics: Vec<Diagnostic> = self.add_diagnostics.clone();
587
588        // Rename overrides must target a type that materialized.
589        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        // Extraction failures: hard errors, or skipped with a warning.
598        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        // Validate type references.
618        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                // Transitively omit types referencing skipped or missing types.
639                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        // The set of types actually emitted, in stable order.
674        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        // Import conflicts across everything emitted.
683        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            // Through the shared path so it dedups and conflict-checks like
692            // any user import.
693            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        // Topological sort (Kahn's) so dependencies emit before dependents;
708        // ties resolve in BTreeMap (name) order.
709        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        // With JIT enabled, cross-references resolve to the raw `$` codecs so
717        // every export is compiled over the uncompiled interpreter graph.
718        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        // Assemble the output.
728        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        // Cycles (only possible through user-provided Raw/TypeRef loops):
806        // append the remaining names in stable order.
807        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            // The compile functions detect a withFormat-bound codec and
890            // prewarm for the bound format, so the JIT wrap stays outermost.
891            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    /// Generate the bindings and write them to `path`.
908    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        // "AOuter" sorts before "Inner" alphabetically, but references it.
1039        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        // Before add.
1076        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        // After add.
1085        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        // Field order is layout, so it must survive the relabelling.
1151        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        // Variant tags keep their own (default: preserved) casing.
1176        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        // Two names that only collide *after* conversion are fine as-is.
1268        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        // User modules keep their exact specifier.
1369        assert!(code.contains("import { MyCodec } from './my-codec.ts';"));
1370        // Emitted factory calls and type exports are direction-independent.
1371        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        // Each split module is a sibling file of the one it splits, so the
1394        // specifier gains a `.decode`/`.encode` segment — except the package
1395        // root, whose counterpart is the standalone `decode` module.
1396        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        // Type exports still derive from the (drop-in) compiled exports.
1428        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        // The compiled Outer export must see Inner's interpreter codec, not
1439        // the opaque compiled one, so the JIT can inline across types.
1440        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        // withFormat stays inside the compileCodec wrap.
1454        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        // The JIT import is emitted direction-matched, not rewritten — the
1492        // rewrite must not append a second suffix to it.
1493        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}