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