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::error::{Diagnostic, DiagnosticKind, Error, SourceLocation};
8use crate::expr::{CodecExpr, generate_import_block};
9use crate::registry::{ExternalType, Registry, WithWrapper};
10
11/// How to handle a field whose type cannot be mapped to a codec.
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
13pub enum OnUnknown {
14    /// Aggregate a diagnostic and fail [`generate`](CodeGenerator::generate).
15    #[default]
16    Error,
17    /// Emit a `cargo:warning` and omit the containing type — and,
18    /// transitively, every type referencing it — from the output.
19    SkipContainingType,
20}
21
22/// An enum variant for [`CodeGenerator::add_enum`].
23#[derive(Debug, Clone)]
24pub enum EnumVariant {
25    /// A unit variant: `Name` - emitted as `Name: null`.
26    Unit(String),
27    /// A newtype (1-tuple) variant: `Name(T)` - emitted as a bare codec.
28    Newtype(String, CodecExpr),
29    /// An n-tuple variant (n >= 2): `Name(T0, T1)` - emitted as an array of codecs (`[t0, t1]`), decoded as an array value.
30    /// The fields stay flattened in the enum layout (this is NOT a nested `r.tuple` block).
31    Tuple(String, Vec<CodecExpr>),
32    /// A struct variant: `Name { a: T }` — emitted as a record of codecs.
33    Struct(String, Vec<(String, CodecExpr)>),
34}
35
36impl EnumVariant {
37    /// The variant name.
38    pub fn name(&self) -> &str {
39        match self {
40            EnumVariant::Unit(name)
41            | EnumVariant::Newtype(name, _)
42            | EnumVariant::Tuple(name, _)
43            | EnumVariant::Struct(name, _) => name,
44        }
45    }
46}
47
48/// The kind-specific payload of a generated type.
49#[derive(Debug, Clone)]
50pub(crate) enum TypeKind {
51    Struct(Vec<(String, CodecExpr)>),
52    Enum(Vec<EnumVariant>),
53    Alias(CodecExpr),
54}
55
56/// The non-default wire format configured via [`set_format`](CodeGenerator::set_format).
57#[derive(Debug, Clone)]
58struct FormatSpec {
59    endian: String,
60    pointer_width: u32,
61    aligned: bool,
62}
63
64impl FormatSpec {
65    fn is_default(&self) -> bool {
66        self.endian == "little" && self.pointer_width == 32 && self.aligned
67    }
68
69    /// The non-default keys as `r.format(...)` options.
70    fn options(&self) -> String {
71        let mut entries = Vec::new();
72        if self.endian != "little" {
73            entries.push(format!("endian: '{}'", self.endian));
74        }
75        if self.pointer_width != 32 {
76            entries.push(format!("pointerWidth: {}", self.pointer_width));
77        }
78        if !self.aligned {
79            entries.push("aligned: false".to_string());
80        }
81        entries.join(", ")
82    }
83}
84
85/// Collects type definitions — from Rust sources or programmatically — and
86/// generates TypeScript codec bindings for the `rkyv-js` runtime.
87///
88/// # Example
89///
90/// ```
91/// use rkyv_js_codegen::{CodeGenerator, codec};
92///
93/// let mut generator = CodeGenerator::new();
94/// generator.add_struct("Point", [("x", codec::f64()), ("y", codec::f64())]);
95/// let code = generator.generate().unwrap();
96/// assert!(code.contains("export const ArchivedPoint = r.struct({"));
97/// assert!(code.contains("export type Point = r.Infer<typeof ArchivedPoint>;"));
98/// ```
99#[derive(Debug)]
100pub struct CodeGenerator {
101    /// Successfully added types, keyed by Rust type name.
102    pub(crate) types: BTreeMap<String, TypeKind>,
103    /// Types whose extraction produced diagnostics, keyed by Rust type name.
104    pub(crate) failed: BTreeMap<String, Vec<Diagnostic>>,
105    /// Diagnostics recorded at add time (duplicate type names).
106    pub(crate) add_diagnostics: Vec<Diagnostic>,
107    /// `set_archived_name` overrides, applied at generate time.
108    overrides: BTreeMap<String, String>,
109    header: Option<String>,
110    allow_typescript_syntax: bool,
111    pub(crate) on_unknown: OnUnknown,
112    /// Derive paths that mark a type for extraction.
113    pub(crate) marker_paths: BTreeSet<String>,
114    pub(crate) registry: Registry,
115    format: Option<FormatSpec>,
116    direction: Direction,
117    jit: bool,
118}
119
120/// Which half of the codec surface the generated bindings target.
121///
122/// The emitted factory calls and type exports are identical in all three modes.
123/// Only the `rkyv-js` import specifiers change, so a decode-only bundle never pulls the writer/hasher machinery (and vice versa).
124#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
125pub enum Direction {
126    /// Full codecs (`rkyv-js`): encode + decode + access.
127    #[default]
128    Full,
129    /// Decoder-only bindings (`rkyv-js/decode`, `rkyv-js/lib/*/decode`).
130    Decode,
131    /// Encoder-only bindings (`rkyv-js/encode`, `rkyv-js/lib/*/encode`).
132    Encode,
133}
134
135impl Direction {
136    fn suffix(self) -> Option<&'static str> {
137        match self {
138            Direction::Full => None,
139            Direction::Decode => Some("/decode"),
140            Direction::Encode => Some("/encode"),
141        }
142    }
143
144    /// The JIT entry point and compile function for this direction.
145    fn jit_entry(self) -> (&'static str, &'static str) {
146        match self {
147            Direction::Full => ("rkyv-js/jit", "compileCodec"),
148            Direction::Decode => ("rkyv-js/jit/decode", "compileDecoder"),
149            Direction::Encode => ("rkyv-js/jit/encode", "compileEncoder"),
150        }
151    }
152
153    /// Rewrite an emitted import block's `rkyv-js` specifiers for this direction.
154    /// Non-`rkyv-js` specifiers (user `register_external` modules) are left untouched.
155    /// Hand-written codecs must provide their own direction-appropriate exports.
156    pub(crate) fn rewrite_import_block(self, block: &str) -> String {
157        let Some(suffix) = self.suffix() else {
158            return block.to_string();
159        };
160        let mut out = String::with_capacity(block.len() + 64);
161        for line in block.lines() {
162            if let Some(spec_start) = line.rfind(" from '").map(|i| i + " from '".len())
163                && let Some(len) = line[spec_start..].find('\'')
164            {
165                let spec = &line[spec_start..spec_start + len];
166                if spec == "rkyv-js" || spec.starts_with("rkyv-js/lib/") {
167                    out.push_str(&line[..spec_start + len]);
168                    out.push_str(suffix);
169                    out.push_str(&line[spec_start + len..]);
170                    out.push('\n');
171                    continue;
172                }
173            }
174            out.push_str(line);
175            out.push('\n');
176        }
177        out
178    }
179}
180
181impl Default for CodeGenerator {
182    fn default() -> Self {
183        Self::new()
184    }
185}
186
187impl CodeGenerator {
188    /// Create a generator with the built-in type and wrapper registrations.
189    pub fn new() -> Self {
190        Self {
191            types: BTreeMap::new(),
192            failed: BTreeMap::new(),
193            add_diagnostics: Vec::new(),
194            overrides: BTreeMap::new(),
195            header: None,
196            allow_typescript_syntax: true,
197            on_unknown: OnUnknown::Error,
198            marker_paths: BTreeSet::from(["rkyv::Archive".to_string()]),
199            registry: Registry::with_builtins(),
200            format: None,
201            direction: Direction::Full,
202            jit: false,
203        }
204    }
205
206    /// Replace the header comment of the generated file.
207    pub fn set_header(&mut self, header: impl Into<String>) -> &mut Self {
208        self.header = Some(header.into());
209        self
210    }
211
212    /// Emit unidirectional bindings: [`Direction::Decode`] rewrites every `rkyv-js` import specifier
213    /// to its `/decode` counterpart (`rkyv-js/lib/X` becomes `rkyv-js/lib/X/decode`), [`Direction::Encode`] symmetrically.
214    ///
215    /// Factory names and type exports are unchanged;
216    /// imports of user modules registered via `register_external` are not rewritten.
217    pub fn set_direction(&mut self, direction: Direction) -> &mut Self {
218        self.direction = direction;
219        self
220    }
221
222    /// Wrap every exported codec in the direction-matched JIT compile function: `compileCodec` from `rkyv-js/jit` for [`Direction::Full`],
223    /// `compileDecoder` from `rkyv-js/jit/decode` resp. `compileEncoder` from `rkyv-js/jit/encode` for unidirectional bindings.
224    ///
225    /// Each type is emitted as a non-exported interpreter codec (`const {Name}$ = ...`)
226    /// plus a compiled export (`export const {Name} = compileCodec({Name}$);`),
227    /// and cross-references between generated types resolve to the `$` codecs:
228    /// a compiled codec is opaque to the JIT, so compiling each export over the
229    /// raw graph is what lets nested types inline instead of degrading to
230    /// per-element dispatch calls. The compiled exports stay drop-in
231    /// (`encode`/`decode`/`access`/... and `r.Infer` are unchanged),
232    /// and fall back to the interpreter codec where `new Function` is blocked (CSP).
233    ///
234    /// Every export compiles eagerly at module load.
235    ///
236    /// Defaults to `false`.
237    pub fn set_jit(&mut self, enabled: bool) -> &mut Self {
238        self.jit = enabled;
239        self
240    }
241
242    /// When `false`, `export type ... = r.Infer<...>` lines are dropped so the output is valid plain JavaScript.
243    ///
244    /// Defaults to `true`.
245    pub fn allow_typescript_syntax(&mut self, enabled: bool) -> &mut Self {
246        self.allow_typescript_syntax = enabled;
247        self
248    }
249
250    /// Configure how unmappable field types are handled.
251    ///
252    /// Defaults to [`OnUnknown::Error`].
253    pub fn on_unknown_type(&mut self, mode: OnUnknown) -> &mut Self {
254        self.on_unknown = mode;
255        self
256    }
257
258    /// Register an additional derive path that marks types for extraction,
259    /// alongside the default `rkyv::Archive`.
260    pub fn add_marker_path(&mut self, path: impl Into<String>) -> &mut Self {
261        self.marker_paths.insert(path.into());
262        self
263    }
264
265    /// Register (or replace) an external type mapping for a fully-qualified Rust path.
266    ///
267    /// ```
268    /// use rkyv_js_codegen::{CodeGenerator, CodecExpr, ExternalType};
269    ///
270    /// let mut generator = CodeGenerator::new();
271    /// generator.register_external(
272    ///     "my_crate::MyVec",
273    ///     ExternalType::generic1(|t| {
274    ///         CodecExpr::call(CodecExpr::import_from("my-pkg/codecs", "myVec"), [t])
275    ///     }),
276    /// );
277    /// ```
278    pub fn register_external(
279        &mut self,
280        path: impl Into<String>,
281        external: ExternalType,
282    ) -> &mut Self {
283        self.registry.register_type(path, external);
284        self
285    }
286
287    /// Register (or replace) a `#[rkyv(with = ...)]` wrapper handler.
288    pub fn register_with(&mut self, path: impl Into<String>, wrapper: WithWrapper) -> &mut Self {
289        self.registry.register_wrapper(path, wrapper);
290        self
291    }
292
293    /// Remove an external type mapping (e.g. to disable a builtin).
294    pub fn unregister_external(&mut self, path: &str) -> &mut Self {
295        self.registry.unregister_type(path);
296        self
297    }
298
299    /// Add a struct definition.
300    pub fn add_struct(
301        &mut self,
302        name: impl Into<String>,
303        fields: impl IntoIterator<Item = (impl Into<String>, CodecExpr)>,
304    ) -> &mut Self {
305        let fields = fields
306            .into_iter()
307            .map(|(field, expr)| (field.into(), expr))
308            .collect();
309        self.add_type(name.into(), TypeKind::Struct(fields), None);
310        self
311    }
312
313    /// Add an enum definition.
314    pub fn add_enum(
315        &mut self,
316        name: impl Into<String>,
317        variants: impl IntoIterator<Item = EnumVariant>,
318    ) -> &mut Self {
319        let variants = variants.into_iter().collect();
320        self.add_type(name.into(), TypeKind::Enum(variants), None);
321        self
322    }
323
324    /// Add a type alias: `export const Archived{name} = <expr>;`.
325    pub fn add_alias(&mut self, name: impl Into<String>, target: CodecExpr) -> &mut Self {
326        self.add_type(name.into(), TypeKind::Alias(target), None);
327        self
328    }
329
330    /// Record a type entry, diagnosing duplicate names.
331    pub(crate) fn add_type(
332        &mut self,
333        name: String,
334        kind: TypeKind,
335        location: Option<SourceLocation>,
336    ) {
337        if self.is_known_type(&name) {
338            self.add_diagnostics.push(
339                Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
340            );
341            return;
342        }
343        self.types.insert(name, kind);
344    }
345
346    /// Record a type whose extraction produced diagnostics.
347    pub(crate) fn add_failed_type(
348        &mut self,
349        name: String,
350        diagnostics: Vec<Diagnostic>,
351        location: Option<SourceLocation>,
352    ) {
353        if self.is_known_type(&name) {
354            self.add_diagnostics.push(
355                Diagnostic::new(DiagnosticKind::DuplicateType { name }).at(location),
356            );
357            return;
358        }
359        self.failed.insert(name, diagnostics);
360    }
361
362    fn is_known_type(&self, name: &str) -> bool {
363        self.types.contains_key(name) || self.failed.contains_key(name)
364    }
365
366    /// Override the archived (exported) name of a type, corresponding to `#[rkyv(archived = Name)]`.
367    ///
368    /// Order-independent: the target type may be added before or after this call.
369    /// A target that never materializes is reported as [`DiagnosticKind::UnknownRenameTarget`] at generate time.
370    pub fn set_archived_name(
371        &mut self,
372        type_name: impl Into<String>,
373        archived_name: impl Into<String>,
374    ) -> &mut Self {
375        self.overrides.insert(type_name.into(), archived_name.into());
376        self
377    }
378
379    /// The archived (exported) name a type will be emitted under, or `None`
380    /// if no type with that name has been added.
381    pub fn archived_name_of(&self, type_name: &str) -> Option<String> {
382        if !self.is_known_type(type_name) {
383            return None;
384        }
385        Some(self.resolved_archived_name(type_name))
386    }
387
388    fn resolved_archived_name(&self, type_name: &str) -> String {
389        self.overrides
390            .get(type_name)
391            .cloned()
392            .unwrap_or_else(|| format!("Archived{type_name}"))
393    }
394
395    /// Configure the rkyv wire format of the generated bindings.
396    ///
397    /// When the format differs from the default (`little`/32/aligned),
398    /// the output declares `const FORMAT = r.format({ ... })` with the non-default keys
399    /// and wraps every exported codec in `r.withFormat(<expr>, FORMAT)`.
400    pub fn set_format(&mut self, endian: &str, pointer_width: u32, aligned: bool) -> &mut Self {
401        self.format = Some(FormatSpec {
402            endian: endian.to_string(),
403            pointer_width,
404            aligned,
405        });
406        self
407    }
408
409    /// The active non-default format, if any.
410    fn nondefault_format(&self) -> Option<&FormatSpec> {
411        self.format.as_ref().filter(|spec| !spec.is_default())
412    }
413
414    /// Every codec expression of a type, labelled with its `Type.field`
415    /// provenance for diagnostics.
416    fn exprs_with_context<'a>(
417        type_name: &str,
418        kind: &'a TypeKind,
419    ) -> Vec<(String, &'a CodecExpr)> {
420        match kind {
421            TypeKind::Struct(fields) => fields
422                .iter()
423                .map(|(field, expr)| (format!("{type_name}.{field}"), expr))
424                .collect(),
425            TypeKind::Enum(variants) => {
426                let mut out = Vec::new();
427                for variant in variants {
428                    match variant {
429                        EnumVariant::Unit(_) => {}
430                        EnumVariant::Newtype(vname, expr) => {
431                            out.push((format!("{type_name}::{vname}"), expr));
432                        }
433                        EnumVariant::Tuple(vname, exprs) => {
434                            for (i, expr) in exprs.iter().enumerate() {
435                                out.push((format!("{type_name}::{vname}.{i}"), expr));
436                            }
437                        }
438                        EnumVariant::Struct(vname, fields) => {
439                            for (field, expr) in fields {
440                                out.push((format!("{type_name}::{vname}.{field}"), expr));
441                            }
442                        }
443                    }
444                }
445                out
446            }
447            TypeKind::Alias(expr) => vec![(type_name.to_string(), expr)],
448        }
449    }
450
451    /// Generate the TypeScript bindings.
452    ///
453    /// Validation runs first; every problem is aggregated into a single [`Error::Codegen`].
454    pub fn generate(&self) -> Result<String, Error> {
455        let mut diagnostics: Vec<Diagnostic> = self.add_diagnostics.clone();
456
457        // Rename overrides must target a type that materialized.
458        for target in self.overrides.keys() {
459            if !self.is_known_type(target) {
460                diagnostics.push(Diagnostic::new(DiagnosticKind::UnknownRenameTarget {
461                    type_name: target.clone(),
462                }));
463            }
464        }
465
466        // Extraction failures: hard errors, or skipped with a warning.
467        let mut skipped: BTreeSet<String> = BTreeSet::new();
468        match self.on_unknown {
469            OnUnknown::Error => {
470                for failure_diagnostics in self.failed.values() {
471                    diagnostics.extend(failure_diagnostics.iter().cloned());
472                }
473            }
474            OnUnknown::SkipContainingType => {
475                for (name, failure_diagnostics) in &self.failed {
476                    skipped.insert(name.clone());
477                    for diagnostic in failure_diagnostics {
478                        eprintln!(
479                            "cargo:warning=rkyv-js-codegen: skipping `{name}`: {diagnostic}"
480                        );
481                    }
482                }
483            }
484        }
485
486        // Validate type references.
487        match self.on_unknown {
488            OnUnknown::Error => {
489                for (name, kind) in &self.types {
490                    for (context, expr) in Self::exprs_with_context(name, kind) {
491                        let mut refs = BTreeSet::new();
492                        expr.collect_type_refs(&mut refs);
493                        for reference in refs {
494                            if !self.is_known_type(&reference) {
495                                diagnostics.push(
496                                    Diagnostic::new(DiagnosticKind::UnresolvedTypeRef {
497                                        name: reference,
498                                    })
499                                    .referenced_by(context.clone()),
500                                );
501                            }
502                        }
503                    }
504                }
505            }
506            OnUnknown::SkipContainingType => {
507                // Transitively omit types referencing skipped or missing types.
508                loop {
509                    let mut newly_skipped = Vec::new();
510                    for (name, kind) in &self.types {
511                        if skipped.contains(name) {
512                            continue;
513                        }
514                        let broken = Self::exprs_with_context(name, kind).iter().any(
515                            |(_, expr)| {
516                                let mut refs = BTreeSet::new();
517                                expr.collect_type_refs(&mut refs);
518                                refs.iter().any(|reference| {
519                                    skipped.contains(reference)
520                                        || !self.types.contains_key(reference)
521                                })
522                            },
523                        );
524                        if broken {
525                            newly_skipped.push(name.clone());
526                        }
527                    }
528                    if newly_skipped.is_empty() {
529                        break;
530                    }
531                    for name in newly_skipped {
532                        eprintln!(
533                            "cargo:warning=rkyv-js-codegen: skipping `{name}`: it references \
534                             a type that was omitted or never added"
535                        );
536                        skipped.insert(name);
537                    }
538                }
539            }
540        }
541
542        // The set of types actually emitted, in stable order.
543        let emitted: BTreeMap<&String, &TypeKind> = self
544            .types
545            .iter()
546            .filter(|(name, _)| !skipped.contains(*name))
547            .collect();
548
549        // Import conflicts across everything emitted.
550        let (jit_module, jit_fn) = self.direction.jit_entry();
551        let jit_import = CodecExpr::import_from(jit_module, jit_fn);
552        let mut all_exprs: Vec<&CodecExpr> = emitted
553            .iter()
554            .flat_map(|(name, kind)| Self::exprs_with_context(name, kind))
555            .map(|(_, expr)| expr)
556            .collect();
557        if self.jit && !emitted.is_empty() {
558            // Through the shared path so it dedups and conflict-checks like
559            // any user import.
560            all_exprs.push(&jit_import);
561        }
562        let import_block = match generate_import_block(all_exprs.iter().copied()) {
563            Ok(block) => self.direction.rewrite_import_block(&block),
564            Err(conflicts) => {
565                diagnostics.extend(conflicts.into_iter().map(Diagnostic::new));
566                String::new()
567            }
568        };
569
570        if !diagnostics.is_empty() {
571            return Err(Error::Codegen(diagnostics));
572        }
573
574        // Topological sort (Kahn's) so dependencies emit before dependents;
575        // ties resolve in BTreeMap (name) order.
576        let order = Self::topological_sort(&emitted);
577
578        let archived_names: BTreeMap<String, String> = emitted
579            .keys()
580            .map(|name| ((*name).clone(), self.resolved_archived_name(name)))
581            .collect();
582
583        // With JIT enabled, cross-references resolve to the raw `$` codecs so
584        // every export is compiled over the uncompiled interpreter graph.
585        let codec_names: BTreeMap<String, String> = if self.jit {
586            archived_names
587                .iter()
588                .map(|(name, archived)| (name.clone(), format!("{archived}$")))
589                .collect()
590        } else {
591            archived_names.clone()
592        };
593
594        // Assemble the output.
595        let mut blocks: Vec<String> = Vec::new();
596
597        let header = self
598            .header
599            .as_deref()
600            .unwrap_or("Auto-generated by rkyv-js-codegen\nDO NOT EDIT MANUALLY");
601        let mut header_block = String::from("/**\n");
602        for line in header.lines() {
603            if line.is_empty() {
604                header_block.push_str(" *\n");
605            } else {
606                header_block.push_str(" * ");
607                header_block.push_str(line);
608                header_block.push('\n');
609            }
610        }
611        header_block.push_str(" */");
612        blocks.push(header_block);
613
614        blocks.push(import_block.trim_end().to_string());
615
616        if let Some(spec) = self.nondefault_format() {
617            blocks.push(format!("const FORMAT = r.format({{ {} }});", spec.options()));
618        }
619
620        for name in &order {
621            let kind = emitted.get(name).expect("ordered names come from emitted");
622            blocks.push(self.emit_type(name, kind, &archived_names, &codec_names));
623        }
624
625        Ok(blocks.join("\n\n") + "\n")
626    }
627
628    fn topological_sort(emitted: &BTreeMap<&String, &TypeKind>) -> Vec<String> {
629        let mut deps: BTreeMap<&str, BTreeSet<String>> = BTreeMap::new();
630        for (name, kind) in emitted {
631            let mut refs = BTreeSet::new();
632            for (_, expr) in Self::exprs_with_context(name, kind) {
633                expr.collect_type_refs(&mut refs);
634            }
635            refs.retain(|reference| {
636                emitted.contains_key(reference) && reference != name.as_str()
637            });
638            deps.insert(name.as_str(), refs);
639        }
640
641        let mut dependents: BTreeMap<&str, Vec<&str>> = BTreeMap::new();
642        let mut in_degree: BTreeMap<&str, usize> = BTreeMap::new();
643        for (name, type_deps) in &deps {
644            in_degree.insert(name, type_deps.len());
645            for dep in type_deps {
646                dependents.entry(dep.as_str()).or_default().push(name);
647            }
648        }
649
650        let mut ready: BTreeSet<&str> = in_degree
651            .iter()
652            .filter(|(_, degree)| **degree == 0)
653            .map(|(name, _)| *name)
654            .collect();
655        let mut order: Vec<String> = Vec::new();
656        let mut done: BTreeSet<&str> = BTreeSet::new();
657
658        while let Some(name) = ready.pop_first() {
659            order.push(name.to_string());
660            done.insert(name);
661            if let Some(children) = dependents.get(name) {
662                for child in children {
663                    let degree = in_degree.get_mut(child).unwrap();
664                    *degree -= 1;
665                    if *degree == 0 {
666                        ready.insert(child);
667                    }
668                }
669            }
670        }
671
672        // Cycles (only possible through user-provided Raw/TypeRef loops):
673        // append the remaining names in stable order.
674        for name in deps.keys() {
675            if !done.contains(name) {
676                order.push((*name).to_string());
677            }
678        }
679
680        order
681    }
682
683    fn emit_type(
684        &self,
685        name: &str,
686        kind: &TypeKind,
687        archived_names: &BTreeMap<String, String>,
688        codec_names: &BTreeMap<String, String>,
689    ) -> String {
690        let archived = archived_names
691            .get(name)
692            .expect("emitted types have archived names")
693            .clone();
694        let render = |expr: &CodecExpr| -> String {
695            expr.render(codec_names)
696                .expect("type references are validated before emission")
697        };
698
699        let codec_expr = match kind {
700            TypeKind::Struct(fields) => {
701                if fields.is_empty() {
702                    "r.struct({})".to_string()
703                } else {
704                    let mut body = String::from("r.struct({\n");
705                    for (field, expr) in fields {
706                        body.push_str(&format!("  {}: {},\n", field, render(expr)));
707                    }
708                    body.push_str("})");
709                    body
710                }
711            }
712            TypeKind::Enum(variants) => {
713                if variants.is_empty() {
714                    "r.taggedEnum({})".to_string()
715                } else {
716                    let mut body = String::from("r.taggedEnum({\n");
717                    for variant in variants {
718                        let value = match variant {
719                            EnumVariant::Unit(_) => "null".to_string(),
720                            EnumVariant::Newtype(_, expr) => render(expr),
721                            EnumVariant::Tuple(_, exprs) => {
722                                render(&CodecExpr::array(exprs.iter().cloned()))
723                            }
724                            EnumVariant::Struct(_, fields) => {
725                                let record = CodecExpr::object(fields.iter().cloned());
726                                render(&record)
727                            }
728                        };
729                        body.push_str(&format!("  {}: {},\n", variant.name(), value));
730                    }
731                    body.push_str("})");
732                    body
733                }
734            }
735            TypeKind::Alias(expr) => render(expr),
736        };
737
738        let codec_expr = match self.nondefault_format() {
739            Some(_) => format!("r.withFormat({codec_expr}, FORMAT)"),
740            None => codec_expr,
741        };
742
743        let mut block = if self.jit {
744            // The compile functions detect a withFormat-bound codec and
745            // prewarm for the bound format, so the JIT wrap stays outermost.
746            let jit_fn = self.direction.jit_entry().1;
747            format!(
748                "const {archived}$ = {codec_expr};\n\n\
749                 export const {archived} = {jit_fn}({archived}$);"
750            )
751        } else {
752            format!("export const {archived} = {codec_expr};")
753        };
754        if self.allow_typescript_syntax {
755            block.push_str(&format!(
756                "\n\nexport type {name} = r.Infer<typeof {archived}>;"
757            ));
758        }
759        block
760    }
761
762    /// Generate the bindings and write them to `path`.
763    pub fn write_to_file(&self, path: impl AsRef<Path>) -> Result<(), Error> {
764        let code = self.generate()?;
765        fs::write(path, code)?;
766        Ok(())
767    }
768}
769
770#[cfg(test)]
771mod tests {
772    use super::*;
773    use crate::expr::codec;
774
775    fn diagnostics(error: Error) -> Vec<Diagnostic> {
776        match error {
777            Error::Codegen(diagnostics) => diagnostics,
778            other => panic!("expected Error::Codegen, got {other:?}"),
779        }
780    }
781
782    #[test]
783    fn struct_emission_snapshot() {
784        let mut generator = CodeGenerator::new();
785        generator.add_struct("Point", [("x", codec::f64()), ("y", codec::f64())]);
786        let code = generator.generate().unwrap();
787        assert_eq!(
788            code,
789            "/**\n\
790             \x20* Auto-generated by rkyv-js-codegen\n\
791             \x20* DO NOT EDIT MANUALLY\n\
792             \x20*/\n\
793             \n\
794             import * as r from 'rkyv-js';\n\
795             \n\
796             export const ArchivedPoint = r.struct({\n\
797             \x20 x: r.f64,\n\
798             \x20 y: r.f64,\n\
799             });\n\
800             \n\
801             export type Point = r.Infer<typeof ArchivedPoint>;\n"
802        );
803    }
804
805    #[test]
806    fn enum_emission_snapshot() {
807        let mut generator = CodeGenerator::new();
808        generator.add_enum(
809            "MixedAlign",
810            [
811                EnumVariant::Struct(
812                    "V".to_string(),
813                    vec![("a".to_string(), codec::u8()), ("b".to_string(), codec::u32())],
814                ),
815                EnumVariant::Newtype("X".to_string(), codec::u64()),
816                EnumVariant::Tuple("Color".to_string(), vec![codec::u8(), codec::u8()]),
817                EnumVariant::Unit("Y".to_string()),
818            ],
819        );
820        let code = generator.generate().unwrap();
821        assert!(code.contains(
822            "export const ArchivedMixedAlign = r.taggedEnum({\n\
823             \x20 V: { a: r.u8, b: r.u32 },\n\
824             \x20 X: r.u64,\n\
825             \x20 Color: [r.u8, r.u8],\n\
826             \x20 Y: null,\n\
827             });"
828        ));
829        assert!(code.contains("export type MixedAlign = r.Infer<typeof ArchivedMixedAlign>;"));
830    }
831
832    #[test]
833    fn alias_emission_snapshot() {
834        let mut generator = CodeGenerator::new();
835        generator.add_alias("UserId", codec::u32());
836        let code = generator.generate().unwrap();
837        assert!(code.contains("export const ArchivedUserId = r.u32;"));
838        assert!(code.contains("export type UserId = r.Infer<typeof ArchivedUserId>;"));
839    }
840
841    #[test]
842    fn imports_are_collected_and_deduped() {
843        let mut generator = CodeGenerator::new();
844        generator.add_struct(
845            "A",
846            [
847                (
848                    "m",
849                    CodecExpr::call(
850                        CodecExpr::import_from("rkyv-js/lib/hashmap", "hashMap"),
851                        [codec::string(), codec::u32()],
852                    ),
853                ),
854                (
855                    "s",
856                    CodecExpr::call(
857                        CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
858                        [codec::string()],
859                    ),
860                ),
861            ],
862        );
863        generator.add_struct(
864            "B",
865            [(
866                "s2",
867                CodecExpr::call(
868                    CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
869                    [codec::u32()],
870                ),
871            )],
872        );
873        let code = generator.generate().unwrap();
874        assert!(code.contains("import { hashMap, hashSet } from 'rkyv-js/lib/hashmap';"));
875        assert_eq!(code.matches("hashSet }").count(), 1);
876    }
877
878    #[test]
879    fn import_conflict_is_reported() {
880        let mut generator = CodeGenerator::new();
881        generator.add_struct("A", [("x", CodecExpr::import_from("pkg-a", "codec"))]);
882        generator.add_struct("B", [("y", CodecExpr::import_from("pkg-b", "codec"))]);
883        let errors = diagnostics(generator.generate().unwrap_err());
884        assert!(errors.iter().any(|diagnostic| matches!(
885            &diagnostic.kind,
886            DiagnosticKind::ImportConflict { export, .. } if export == "codec"
887        )));
888    }
889
890    #[test]
891    fn topo_sort_handles_forward_references() {
892        let mut generator = CodeGenerator::new();
893        // "AOuter" sorts before "Inner" alphabetically, but references it.
894        generator.add_struct("AOuter", [("inner", codec::named("Inner"))]);
895        generator.add_struct("Inner", [("value", codec::u32())]);
896        let code = generator.generate().unwrap();
897        let inner_pos = code.find("export const ArchivedInner").unwrap();
898        let outer_pos = code.find("export const ArchivedAOuter").unwrap();
899        assert!(inner_pos < outer_pos, "dependency must be emitted first");
900        assert!(code.contains("inner: ArchivedInner,"));
901    }
902
903    #[test]
904    fn unresolved_type_ref_reports_referrer() {
905        let mut generator = CodeGenerator::new();
906        generator.add_struct("Outer", [("inner", codec::named("Missing"))]);
907        let errors = diagnostics(generator.generate().unwrap_err());
908        assert_eq!(errors.len(), 1);
909        assert!(matches!(
910            &errors[0].kind,
911            DiagnosticKind::UnresolvedTypeRef { name } if name == "Missing"
912        ));
913        assert_eq!(errors[0].referenced_by.as_deref(), Some("Outer.inner"));
914    }
915
916    #[test]
917    fn duplicate_type_is_reported_at_generate() {
918        let mut generator = CodeGenerator::new();
919        generator.add_struct("Point", [("x", codec::f64())]);
920        generator.add_struct("Point", [("y", codec::f64())]);
921        let errors = diagnostics(generator.generate().unwrap_err());
922        assert!(errors.iter().any(|diagnostic| matches!(
923            &diagnostic.kind,
924            DiagnosticKind::DuplicateType { name } if name == "Point"
925        )));
926    }
927
928    #[test]
929    fn set_archived_name_is_order_independent() {
930        // Before add.
931        let mut generator = CodeGenerator::new();
932        generator.set_archived_name("Foo", "MyFoo");
933        generator.add_struct("Foo", [("x", codec::u32())]);
934        let code = generator.generate().unwrap();
935        assert!(code.contains("export const MyFoo = r.struct({"));
936        assert!(code.contains("export type Foo = r.Infer<typeof MyFoo>;"));
937        assert!(!code.contains("ArchivedFoo"));
938
939        // After add.
940        let mut generator = CodeGenerator::new();
941        generator.add_struct("Foo", [("x", codec::u32())]);
942        generator.set_archived_name("Foo", "MyFoo");
943        let code = generator.generate().unwrap();
944        assert!(code.contains("export const MyFoo = r.struct({"));
945    }
946
947    #[test]
948    fn archived_rename_applies_to_cross_references() {
949        let mut generator = CodeGenerator::new();
950        generator.set_archived_name("Inner", "CustomInner");
951        generator.add_struct("Inner", [("value", codec::u32())]);
952        generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
953        let code = generator.generate().unwrap();
954        assert!(code.contains("export const CustomInner = r.struct({"));
955        assert!(code.contains("inner: CustomInner,"));
956        assert!(!code.contains("ArchivedInner"));
957    }
958
959    #[test]
960    fn unknown_rename_target_is_a_diagnostic() {
961        let mut generator = CodeGenerator::new();
962        generator.add_struct("Foo", [("x", codec::u32())]);
963        generator.set_archived_name("Nope", "MyNope");
964        let errors = diagnostics(generator.generate().unwrap_err());
965        assert!(errors.iter().any(|diagnostic| matches!(
966            &diagnostic.kind,
967            DiagnosticKind::UnknownRenameTarget { type_name } if type_name == "Nope"
968        )));
969    }
970
971    #[test]
972    fn archived_name_of_accessor() {
973        let mut generator = CodeGenerator::new();
974        generator.add_struct("Foo", [("x", codec::u32())]);
975        assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("ArchivedFoo"));
976        generator.set_archived_name("Foo", "MyFoo");
977        assert_eq!(generator.archived_name_of("Foo").as_deref(), Some("MyFoo"));
978        assert_eq!(generator.archived_name_of("Bar"), None);
979    }
980
981    #[test]
982    fn js_mode_omits_type_lines() {
983        let mut generator = CodeGenerator::new();
984        generator.allow_typescript_syntax(false);
985        generator.add_struct("Point", [("x", codec::f64())]);
986        generator.add_alias("UserId", codec::u32());
987        let code = generator.generate().unwrap();
988        assert!(code.contains("export const ArchivedPoint = r.struct({"));
989        assert!(code.contains("export const ArchivedUserId = r.u32;"));
990        assert!(!code.contains("export type"));
991        assert!(!code.contains("r.Infer"));
992    }
993
994    #[test]
995    fn set_format_default_is_a_no_op() {
996        let mut generator = CodeGenerator::new();
997        generator.set_format("little", 32, true);
998        generator.add_struct("Point", [("x", codec::f64())]);
999        let code = generator.generate().unwrap();
1000        assert!(!code.contains("FORMAT"));
1001        assert!(!code.contains("withFormat"));
1002    }
1003
1004    #[test]
1005    fn set_format_nondefault_wraps_exports() {
1006        let mut generator = CodeGenerator::new();
1007        generator.set_format("big", 64, false);
1008        generator.add_struct("Point", [("x", codec::f64())]);
1009        generator.add_alias("UserId", codec::u32());
1010        let code = generator.generate().unwrap();
1011        assert!(code.contains(
1012            "const FORMAT = r.format({ endian: 'big', pointerWidth: 64, aligned: false });"
1013        ));
1014        assert!(code.contains("export const ArchivedPoint = r.withFormat(r.struct({\n"));
1015        assert!(code.contains("}), FORMAT);"));
1016        assert!(code.contains("export const ArchivedUserId = r.withFormat(r.u32, FORMAT);"));
1017    }
1018
1019    #[test]
1020    fn set_format_emits_only_nondefault_keys() {
1021        let mut generator = CodeGenerator::new();
1022        generator.set_format("little", 16, true);
1023        generator.add_struct("Point", [("x", codec::f64())]);
1024        let code = generator.generate().unwrap();
1025        assert!(code.contains("const FORMAT = r.format({ pointerWidth: 16 });"));
1026    }
1027
1028    #[test]
1029    fn custom_header_replaces_default() {
1030        let mut generator = CodeGenerator::new();
1031        generator.set_header("Custom header\nsecond line");
1032        generator.add_struct("Point", [("x", codec::f64())]);
1033        let code = generator.generate().unwrap();
1034        assert!(code.starts_with("/**\n * Custom header\n * second line\n */\n"));
1035        assert!(!code.contains("Auto-generated by rkyv-js-codegen"));
1036    }
1037
1038    #[test]
1039    fn set_direction_full_is_a_no_op() {
1040        let mut generator = CodeGenerator::new();
1041        generator.set_direction(Direction::Full);
1042        generator.add_struct("Point", [("x", codec::f64())]);
1043        let code = generator.generate().unwrap();
1044        assert!(code.contains("import * as r from 'rkyv-js';"));
1045    }
1046
1047    #[test]
1048    fn set_direction_rewrites_rkyv_specifiers_only() {
1049        let mut generator = CodeGenerator::new();
1050        generator.set_direction(Direction::Decode);
1051        generator.add_struct(
1052            "Event",
1053            [
1054                ("id", codec::u32()),
1055                (
1056                    "tags",
1057                    CodecExpr::call(
1058                        CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
1059                        [codec::string()],
1060                    ),
1061                ),
1062                (
1063                    "custom",
1064                    CodecExpr::import_from("./my-codec.ts", "MyCodec"),
1065                ),
1066            ],
1067        );
1068        let code = generator.generate().unwrap();
1069        assert!(code.contains("import * as r from 'rkyv-js/decode';"));
1070        assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap/decode';"));
1071        // User modules keep their exact specifier.
1072        assert!(code.contains("import { MyCodec } from './my-codec.ts';"));
1073        // Emitted factory calls and type exports are direction-independent.
1074        assert!(code.contains("export const ArchivedEvent = r.struct({"));
1075        assert!(code.contains("export type Event = r.Infer<typeof ArchivedEvent>;"));
1076    }
1077
1078    #[test]
1079    fn set_direction_encode_uses_encode_suffix() {
1080        let mut generator = CodeGenerator::new();
1081        generator.set_direction(Direction::Encode);
1082        generator.add_struct("Point", [("x", codec::f64())]);
1083        let code = generator.generate().unwrap();
1084        assert!(code.contains("import * as r from 'rkyv-js/encode';"));
1085    }
1086
1087    #[test]
1088    fn set_jit_wraps_exports() {
1089        let mut generator = CodeGenerator::new();
1090        generator.set_jit(true);
1091        generator.add_struct("Point", [("x", codec::f64())]);
1092        generator.add_alias("UserId", codec::u32());
1093        let code = generator.generate().unwrap();
1094        assert!(code.contains("import { compileCodec } from 'rkyv-js/jit';"));
1095        assert!(code.contains("const ArchivedPoint$ = r.struct({\n"));
1096        assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
1097        assert!(code.contains("const ArchivedUserId$ = r.u32;"));
1098        assert!(code.contains("export const ArchivedUserId = compileCodec(ArchivedUserId$);"));
1099        // Type exports still derive from the (drop-in) compiled exports.
1100        assert!(code.contains("export type Point = r.Infer<typeof ArchivedPoint>;"));
1101    }
1102
1103    #[test]
1104    fn set_jit_references_resolve_to_raw_codecs() {
1105        let mut generator = CodeGenerator::new();
1106        generator.set_jit(true);
1107        generator.add_struct("Inner", [("value", codec::u32())]);
1108        generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
1109        let code = generator.generate().unwrap();
1110        // The compiled Outer export must see Inner's interpreter codec, not
1111        // the opaque compiled one, so the JIT can inline across types.
1112        assert!(code.contains("inner: ArchivedInner$,"));
1113        assert!(code.contains("export const ArchivedInner = compileCodec(ArchivedInner$);"));
1114        assert!(code.contains("export const ArchivedOuter = compileCodec(ArchivedOuter$);"));
1115    }
1116
1117    #[test]
1118    fn set_jit_composes_with_format() {
1119        let mut generator = CodeGenerator::new();
1120        generator.set_jit(true);
1121        generator.set_format("big", 64, true);
1122        generator.add_struct("Point", [("x", codec::f64())]);
1123        let code = generator.generate().unwrap();
1124        assert!(code.contains("const FORMAT = r.format({ endian: 'big', pointerWidth: 64 });"));
1125        // withFormat stays inside the compileCodec wrap.
1126        assert!(code.contains("const ArchivedPoint$ = r.withFormat(r.struct({\n"));
1127        assert!(code.contains("export const ArchivedPoint = compileCodec(ArchivedPoint$);"));
1128    }
1129
1130    #[test]
1131    fn set_jit_respects_archived_renames() {
1132        let mut generator = CodeGenerator::new();
1133        generator.set_jit(true);
1134        generator.set_archived_name("Inner", "CustomInner");
1135        generator.add_struct("Inner", [("value", codec::u32())]);
1136        generator.add_struct("Outer", [("inner", codec::named("Inner"))]);
1137        let code = generator.generate().unwrap();
1138        assert!(code.contains("inner: CustomInner$,"));
1139        assert!(code.contains("export const CustomInner = compileCodec(CustomInner$);"));
1140    }
1141
1142    #[test]
1143    fn set_jit_decode_direction_uses_compile_decoder() {
1144        let mut generator = CodeGenerator::new();
1145        generator.set_jit(true);
1146        generator.set_direction(Direction::Decode);
1147        generator.add_struct(
1148            "Event",
1149            [
1150                ("id", codec::u32()),
1151                (
1152                    "tags",
1153                    CodecExpr::call(
1154                        CodecExpr::import_from("rkyv-js/lib/hashmap", "hashSet"),
1155                        [codec::string()],
1156                    ),
1157                ),
1158            ],
1159        );
1160        let code = generator.generate().unwrap();
1161        assert!(code.contains("import * as r from 'rkyv-js/decode';"));
1162        assert!(code.contains("import { hashSet } from 'rkyv-js/lib/hashmap/decode';"));
1163        // The JIT import is emitted direction-matched, not rewritten.
1164        assert!(code.contains("import { compileDecoder } from 'rkyv-js/jit/decode';"));
1165        assert!(code.contains("export const ArchivedEvent = compileDecoder(ArchivedEvent$);"));
1166        assert!(!code.contains("compileCodec"));
1167    }
1168
1169    #[test]
1170    fn set_jit_encode_direction_uses_compile_encoder() {
1171        let mut generator = CodeGenerator::new();
1172        generator.set_jit(true);
1173        generator.set_direction(Direction::Encode);
1174        generator.add_struct("Point", [("x", codec::f64())]);
1175        let code = generator.generate().unwrap();
1176        assert!(code.contains("import * as r from 'rkyv-js/encode';"));
1177        assert!(code.contains("import { compileEncoder } from 'rkyv-js/jit/encode';"));
1178        assert!(code.contains("export const ArchivedPoint = compileEncoder(ArchivedPoint$);"));
1179    }
1180}