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