Skip to main content

wasmtime_test_util/
component_fuzz.rs

1//! This module generates test cases for the Wasmtime component model function APIs,
2//! e.g. `wasmtime::component::func::Func` and `TypedFunc`.
3//!
4//! Each case includes a list of arbitrary interface types to use as parameters, plus another one to use as a
5//! result, and a component which exports a function and imports a function.  The exported function forwards its
6//! parameters to the imported one and forwards the result back to the caller.  This serves to exercise Wasmtime's
7//! lifting and lowering code and verify the values remain intact during both processes.
8
9use arbitrary::{Arbitrary, Unstructured};
10use indexmap::IndexSet;
11use proc_macro2::{Ident, TokenStream};
12use quote::{ToTokens, format_ident, quote};
13use std::borrow::Cow;
14use std::fmt::{self, Debug, Write};
15use std::hash::{Hash, Hasher};
16use std::iter;
17use std::ops::Deref;
18use wasmtime_component_util::{DiscriminantSize, FlagsSize, REALLOC_AND_FREE};
19
20const MAX_FLAT_PARAMS: usize = 16;
21const MAX_FLAT_ASYNC_PARAMS: usize = 4;
22const MAX_FLAT_RESULTS: usize = 1;
23
24/// The name of the imported host function which the generated component will call
25pub const IMPORT_FUNCTION: &str = "echo-import";
26
27/// The name of the exported guest function which the host should call
28pub const EXPORT_FUNCTION: &str = "echo-export";
29
30/// Wasmtime allows up to 100 type depth so limit this to just under that.
31pub const MAX_TYPE_DEPTH: u32 = 90;
32
33macro_rules! uwriteln {
34    ($($arg:tt)*) => {
35        writeln!($($arg)*).unwrap()
36    };
37}
38
39macro_rules! uwrite {
40    ($($arg:tt)*) => {
41        write!($($arg)*).unwrap()
42    };
43}
44
45#[derive(Debug, Copy, Clone, PartialEq, Eq)]
46enum CoreType {
47    I32,
48    I64,
49    F32,
50    F64,
51}
52
53impl CoreType {
54    /// This is the `join` operation specified in [the canonical
55    /// ABI](https://github.com/WebAssembly/component-model/blob/main/design/mvp/CanonicalABI.md#flattening) for
56    /// variant types.
57    fn join(self, other: Self) -> Self {
58        match (self, other) {
59            _ if self == other => self,
60            (Self::I32, Self::F32) | (Self::F32, Self::I32) => Self::I32,
61            _ => Self::I64,
62        }
63    }
64}
65
66impl fmt::Display for CoreType {
67    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
68        match self {
69            Self::I32 => f.write_str("i32"),
70            Self::I64 => f.write_str("i64"),
71            Self::F32 => f.write_str("f32"),
72            Self::F64 => f.write_str("f64"),
73        }
74    }
75}
76
77/// Wraps a `Box<[T]>` and provides an `Arbitrary` implementation that always generates slices of length less than
78/// or equal to the longest tuple for which Wasmtime generates a `ComponentType` impl
79#[derive(Debug, Clone)]
80pub struct VecInRange<T, const L: u32, const H: u32>(Vec<T>);
81
82impl<T, const L: u32, const H: u32> VecInRange<T, L, H> {
83    fn new<'a>(
84        input: &mut Unstructured<'a>,
85        fuel: &mut u32,
86        generate: impl Fn(&mut Unstructured<'a>, &mut u32) -> arbitrary::Result<T>,
87    ) -> arbitrary::Result<Self> {
88        let mut ret = Vec::new();
89        input.arbitrary_loop(Some(L), Some(H), |input| {
90            if *fuel > 0 {
91                *fuel = *fuel - 1;
92                ret.push(generate(input, fuel)?);
93                Ok(std::ops::ControlFlow::Continue(()))
94            } else {
95                Ok(std::ops::ControlFlow::Break(()))
96            }
97        })?;
98        Ok(Self(ret))
99    }
100}
101
102impl<T, const L: u32, const H: u32> Deref for VecInRange<T, L, H> {
103    type Target = [T];
104
105    fn deref(&self) -> &[T] {
106        self.0.deref()
107    }
108}
109
110/// Represents a component model interface type
111#[expect(missing_docs, reason = "self-describing")]
112#[derive(Debug, Clone)]
113pub enum Type {
114    Bool,
115    S8,
116    U8,
117    S16,
118    U16,
119    S32,
120    U32,
121    S64,
122    U64,
123    Float32,
124    Float64,
125    Char,
126    String,
127    List(Box<Type>),
128    Map(Box<Type>, Box<Type>),
129
130    // Give records the ability to generate a generous amount of fields but
131    // don't let the fuzzer go too wild since `wasmparser`'s validator currently
132    // has hard limits in the 1000-ish range on the number of fields a record
133    // may contain.
134    Record(VecInRange<Type, 1, 200>),
135
136    // Tuples can only have up to 16 type parameters in wasmtime right now for
137    // the static API, but the standard library only supports `Debug` up to 11
138    // elements, so compromise at an even 10.
139    Tuple(VecInRange<Type, 1, 10>),
140
141    // Like records, allow a good number of variants, but variants require at
142    // least one case.
143    Variant(VecInRange<Option<Type>, 1, 200>),
144    Enum(u32),
145
146    Option(Box<Type>),
147    Result {
148        ok: Option<Box<Type>>,
149        err: Option<Box<Type>>,
150    },
151
152    Flags(u32),
153}
154
155impl Type {
156    pub fn generate(
157        u: &mut Unstructured<'_>,
158        depth: u32,
159        fuel: &mut u32,
160    ) -> arbitrary::Result<Type> {
161        *fuel = fuel.saturating_sub(1);
162        let max = if depth == 0 || *fuel == 0 { 12 } else { 21 };
163        Ok(match u.int_in_range(0..=max)? {
164            0 => Type::Bool,
165            1 => Type::S8,
166            2 => Type::U8,
167            3 => Type::S16,
168            4 => Type::U16,
169            5 => Type::S32,
170            6 => Type::U32,
171            7 => Type::S64,
172            8 => Type::U64,
173            9 => Type::Float32,
174            10 => Type::Float64,
175            11 => Type::Char,
176            12 => Type::String,
177            // ^-- if you add something here update the `depth == 0` case above
178            13 => Type::List(Box::new(Type::generate(u, depth - 1, fuel)?)),
179            14 => Type::Record(Type::generate_list(u, depth - 1, fuel)?),
180            15 => Type::Tuple(Type::generate_list(u, depth - 1, fuel)?),
181            16 => Type::Variant(VecInRange::new(u, fuel, |u, fuel| {
182                Type::generate_opt(u, depth - 1, fuel)
183            })?),
184            17 => {
185                let amt = u.int_in_range(1..=(*fuel).max(1).min(257))?;
186                *fuel -= amt;
187                Type::Enum(amt)
188            }
189            18 => Type::Option(Box::new(Type::generate(u, depth - 1, fuel)?)),
190            19 => Type::Result {
191                ok: Type::generate_opt(u, depth - 1, fuel)?.map(Box::new),
192                err: Type::generate_opt(u, depth - 1, fuel)?.map(Box::new),
193            },
194            20 => {
195                let amt = u.int_in_range(1..=(*fuel).min(32))?;
196                *fuel -= amt;
197                Type::Flags(amt)
198            }
199            21 => Type::Map(
200                Box::new(Type::generate_hashable_key(u, fuel)?),
201                Box::new(Type::generate(u, depth - 1, fuel)?),
202            ),
203            // ^-- if you add something here update the `depth != 0` case above
204            _ => unreachable!(),
205        })
206    }
207
208    /// Generate a type that can be used as a HashMap key (implements Hash + Eq).
209    /// This excludes floats and complex types that might contain floats.
210    fn generate_hashable_key(u: &mut Unstructured<'_>, fuel: &mut u32) -> arbitrary::Result<Type> {
211        *fuel = fuel.saturating_sub(1);
212        // Only generate types that implement Hash and Eq:
213        // - No Float32/Float64 (NaN comparison issues)
214        // - No complex types (Record, Tuple, Variant, etc.) as they might contain floats
215        // - String is allowed as it implements Hash + Eq
216        Ok(match u.int_in_range(0..=10)? {
217            0 => Type::Bool,
218            1 => Type::S8,
219            2 => Type::U8,
220            3 => Type::S16,
221            4 => Type::U16,
222            5 => Type::S32,
223            6 => Type::U32,
224            7 => Type::S64,
225            8 => Type::U64,
226            9 => Type::Char,
227            10 => Type::String,
228            _ => unreachable!(),
229        })
230    }
231
232    fn generate_opt(
233        u: &mut Unstructured<'_>,
234        depth: u32,
235        fuel: &mut u32,
236    ) -> arbitrary::Result<Option<Type>> {
237        Ok(if u.arbitrary()? {
238            Some(Type::generate(u, depth, fuel)?)
239        } else {
240            None
241        })
242    }
243
244    fn generate_list<const L: u32, const H: u32>(
245        u: &mut Unstructured<'_>,
246        depth: u32,
247        fuel: &mut u32,
248    ) -> arbitrary::Result<VecInRange<Type, L, H>> {
249        VecInRange::new(u, fuel, |u, fuel| Type::generate(u, depth, fuel))
250    }
251
252    /// Generates text format wasm into `s` to store a value of this type, in
253    /// its flat representation stored in the `locals` provided, to the local
254    /// named `ptr` at the `offset` provided.
255    ///
256    /// This will register helper functions necessary in `helpers`. The
257    /// `locals` iterator will be advanced for all locals consumed by this
258    /// store operation.
259    fn store_flat<'a>(
260        &'a self,
261        s: &mut String,
262        ptr: &str,
263        offset: u32,
264        locals: &mut dyn Iterator<Item = FlatSource>,
265        helpers: &mut IndexSet<Helper<'a>>,
266    ) {
267        enum Kind {
268            Primitive(&'static str),
269            PointerPair,
270            Helper,
271        }
272        let kind = match self {
273            Type::Bool | Type::S8 | Type::U8 => Kind::Primitive("i32.store8"),
274            Type::S16 | Type::U16 => Kind::Primitive("i32.store16"),
275            Type::S32 | Type::U32 | Type::Char => Kind::Primitive("i32.store"),
276            Type::S64 | Type::U64 => Kind::Primitive("i64.store"),
277            Type::Float32 => Kind::Primitive("f32.store"),
278            Type::Float64 => Kind::Primitive("f64.store"),
279            Type::String | Type::List(_) | Type::Map(_, _) => Kind::PointerPair,
280            Type::Enum(n) if *n <= (1 << 8) => Kind::Primitive("i32.store8"),
281            Type::Enum(n) if *n <= (1 << 16) => Kind::Primitive("i32.store16"),
282            Type::Enum(_) => Kind::Primitive("i32.store"),
283            Type::Flags(n) if *n <= 8 => Kind::Primitive("i32.store8"),
284            Type::Flags(n) if *n <= 16 => Kind::Primitive("i32.store16"),
285            Type::Flags(n) if *n <= 32 => Kind::Primitive("i32.store"),
286            Type::Flags(_) => unreachable!(),
287            Type::Record(_)
288            | Type::Tuple(_)
289            | Type::Variant(_)
290            | Type::Option(_)
291            | Type::Result { .. } => Kind::Helper,
292        };
293
294        match kind {
295            Kind::Primitive(op) => uwriteln!(
296                s,
297                "({op} offset={offset} (local.get {ptr}) {})",
298                locals.next().unwrap()
299            ),
300            Kind::PointerPair => {
301                let abi_ptr = locals.next().unwrap();
302                let abi_len = locals.next().unwrap();
303                uwriteln!(s, "(i32.store offset={offset} (local.get {ptr}) {abi_ptr})",);
304                let offset = offset + 4;
305                uwriteln!(s, "(i32.store offset={offset} (local.get {ptr}) {abi_len})",);
306            }
307            Kind::Helper => {
308                let (index, _) = helpers.insert_full(Helper(self));
309                uwriteln!(s, "(i32.add (local.get {ptr}) (i32.const {offset}))");
310                for _ in 0..self.lowered().len() {
311                    let i = locals.next().unwrap();
312                    uwriteln!(s, "{i}");
313                }
314                uwriteln!(s, "call $store_helper_{index}");
315            }
316        }
317    }
318
319    /// Generates a text-format wasm function which takes a pointer and this
320    /// type's flat representation as arguments and then stores this value in
321    /// the first argument.
322    ///
323    /// This is used to store records/variants to cut down on the size of final
324    /// functions and make codegen here a bit easier.
325    fn store_flat_helper<'a>(
326        &'a self,
327        s: &mut String,
328        i: usize,
329        helpers: &mut IndexSet<Helper<'a>>,
330    ) {
331        uwrite!(s, "(func $store_helper_{i} (param i32)");
332        let lowered = self.lowered();
333        for ty in &lowered {
334            uwrite!(s, " (param {ty})");
335        }
336        s.push_str("\n");
337        let locals = (0..lowered.len() as u32).map(|i| i + 1).collect::<Vec<_>>();
338        let record = |s: &mut String, helpers: &mut IndexSet<Helper<'a>>, types: &'a [Type]| {
339            let mut locals = locals.iter().cloned().map(FlatSource::Local);
340            for (offset, ty) in record_field_offsets(types) {
341                ty.store_flat(s, "0", offset, &mut locals, helpers);
342            }
343            assert!(locals.next().is_none());
344        };
345        let variant = |s: &mut String,
346                       helpers: &mut IndexSet<Helper<'a>>,
347                       types: &[Option<&'a Type>]| {
348            let (size, offset) = variant_memory_info(types.iter().cloned());
349            // One extra block for out-of-bounds discriminants.
350            for _ in 0..types.len() + 1 {
351                s.push_str("block\n");
352            }
353
354            // Store the discriminant in memory, then branch on it to figure
355            // out which case we're in.
356            let store = match size {
357                DiscriminantSize::Size1 => "i32.store8",
358                DiscriminantSize::Size2 => "i32.store16",
359                DiscriminantSize::Size4 => "i32.store",
360            };
361            uwriteln!(s, "({store} (local.get 0) (local.get 1))");
362            s.push_str("local.get 1\n");
363            s.push_str("br_table");
364            for i in 0..types.len() + 1 {
365                uwrite!(s, " {i}");
366            }
367            s.push_str("\nend\n");
368
369            // Store each payload individually while converting locals from
370            // their source types to the precise type necessary for this
371            // variant.
372            for ty in types {
373                if let Some(ty) = ty {
374                    let ty_lowered = ty.lowered();
375                    let mut locals = locals[1..].iter().zip(&lowered[1..]).zip(&ty_lowered).map(
376                        |((i, from), to)| FlatSource::LocalConvert {
377                            local: *i,
378                            from: *from,
379                            to: *to,
380                        },
381                    );
382                    ty.store_flat(s, "0", offset, &mut locals, helpers);
383                }
384                s.push_str("return\n");
385                s.push_str("end\n");
386            }
387
388            // Catch-all result which is for out-of-bounds discriminants.
389            s.push_str("unreachable\n");
390        };
391        match self {
392            Type::Bool
393            | Type::S8
394            | Type::U8
395            | Type::S16
396            | Type::U16
397            | Type::S32
398            | Type::U32
399            | Type::Char
400            | Type::S64
401            | Type::U64
402            | Type::Float32
403            | Type::Float64
404            | Type::String
405            | Type::List(_)
406            | Type::Map(_, _)
407            | Type::Flags(_)
408            | Type::Enum(_) => unreachable!(),
409
410            Type::Record(r) => record(s, helpers, r),
411            Type::Tuple(t) => record(s, helpers, t),
412            Type::Variant(v) => variant(
413                s,
414                helpers,
415                &v.iter().map(|t| t.as_ref()).collect::<Vec<_>>(),
416            ),
417            Type::Option(o) => variant(s, helpers, &[None, Some(&**o)]),
418            Type::Result { ok, err } => variant(s, helpers, &[ok.as_deref(), err.as_deref()]),
419        };
420        s.push_str(")\n");
421    }
422
423    /// Same as `store_flat`, except loads the flat values from `ptr+offset`.
424    ///
425    /// Results are placed directly on the wasm stack.
426    fn load_flat<'a>(
427        &'a self,
428        s: &mut String,
429        ptr: &str,
430        offset: u32,
431        helpers: &mut IndexSet<Helper<'a>>,
432    ) {
433        enum Kind {
434            Primitive(&'static str),
435            PointerPair,
436            Helper,
437        }
438        let kind = match self {
439            Type::Bool | Type::U8 => Kind::Primitive("i32.load8_u"),
440            Type::S8 => Kind::Primitive("i32.load8_s"),
441            Type::U16 => Kind::Primitive("i32.load16_u"),
442            Type::S16 => Kind::Primitive("i32.load16_s"),
443            Type::U32 | Type::S32 | Type::Char => Kind::Primitive("i32.load"),
444            Type::U64 | Type::S64 => Kind::Primitive("i64.load"),
445            Type::Float32 => Kind::Primitive("f32.load"),
446            Type::Float64 => Kind::Primitive("f64.load"),
447            Type::String | Type::List(_) | Type::Map(_, _) => Kind::PointerPair,
448            Type::Enum(n) if *n <= (1 << 8) => Kind::Primitive("i32.load8_u"),
449            Type::Enum(n) if *n <= (1 << 16) => Kind::Primitive("i32.load16_u"),
450            Type::Enum(_) => Kind::Primitive("i32.load"),
451            Type::Flags(n) if *n <= 8 => Kind::Primitive("i32.load8_u"),
452            Type::Flags(n) if *n <= 16 => Kind::Primitive("i32.load16_u"),
453            Type::Flags(n) if *n <= 32 => Kind::Primitive("i32.load"),
454            Type::Flags(_) => unreachable!(),
455
456            Type::Record(_)
457            | Type::Tuple(_)
458            | Type::Variant(_)
459            | Type::Option(_)
460            | Type::Result { .. } => Kind::Helper,
461        };
462        match kind {
463            Kind::Primitive(op) => uwriteln!(s, "({op} offset={offset} (local.get {ptr}))"),
464            Kind::PointerPair => {
465                uwriteln!(s, "(i32.load offset={offset} (local.get {ptr}))",);
466                let offset = offset + 4;
467                uwriteln!(s, "(i32.load offset={offset} (local.get {ptr}))",);
468            }
469            Kind::Helper => {
470                let (index, _) = helpers.insert_full(Helper(self));
471                uwriteln!(s, "(i32.add (local.get {ptr}) (i32.const {offset}))");
472                uwriteln!(s, "call $load_helper_{index}");
473            }
474        }
475    }
476
477    /// Same as `store_flat_helper` but for loading the flat representation.
478    fn load_flat_helper<'a>(
479        &'a self,
480        s: &mut String,
481        i: usize,
482        helpers: &mut IndexSet<Helper<'a>>,
483    ) {
484        uwrite!(s, "(func $load_helper_{i} (param i32)");
485        let lowered = self.lowered();
486        for ty in &lowered {
487            uwrite!(s, " (result {ty})");
488        }
489        s.push_str("\n");
490        let record = |s: &mut String, helpers: &mut IndexSet<Helper<'a>>, types: &'a [Type]| {
491            for (offset, ty) in record_field_offsets(types) {
492                ty.load_flat(s, "0", offset, helpers);
493            }
494        };
495        let variant = |s: &mut String,
496                       helpers: &mut IndexSet<Helper<'a>>,
497                       types: &[Option<&'a Type>]| {
498            let (size, offset) = variant_memory_info(types.iter().cloned());
499
500            // Destination locals where the flat representation will be stored.
501            // These are automatically zero which handles unused fields too.
502            for (i, ty) in lowered.iter().enumerate() {
503                uwriteln!(s, " (local $r{i} {ty})");
504            }
505
506            // Return block each case jumps to after setting all locals.
507            s.push_str("block $r\n");
508
509            // One extra block for "out of bounds discriminant".
510            for _ in 0..types.len() + 1 {
511                s.push_str("block\n");
512            }
513
514            // Load the discriminant and branch on it, storing it in
515            // `$r0` as well which is the first flat local representation.
516            let load = match size {
517                DiscriminantSize::Size1 => "i32.load8_u",
518                DiscriminantSize::Size2 => "i32.load16",
519                DiscriminantSize::Size4 => "i32.load",
520            };
521            uwriteln!(s, "({load} (local.get 0))");
522            s.push_str("local.tee $r0\n");
523            s.push_str("br_table");
524            for i in 0..types.len() + 1 {
525                uwrite!(s, " {i}");
526            }
527            s.push_str("\nend\n");
528
529            // For each payload, which is in its own block, load payloads from
530            // memory as necessary and convert them into the final locals.
531            for ty in types {
532                if let Some(ty) = ty {
533                    let ty_lowered = ty.lowered();
534                    ty.load_flat(s, "0", offset, helpers);
535                    for (i, (from, to)) in ty_lowered.iter().zip(&lowered[1..]).enumerate().rev() {
536                        let i = i + 1;
537                        match (from, to) {
538                            (CoreType::F32, CoreType::I32) => {
539                                s.push_str("i32.reinterpret_f32\n");
540                            }
541                            (CoreType::I32, CoreType::I64) => {
542                                s.push_str("i64.extend_i32_u\n");
543                            }
544                            (CoreType::F32, CoreType::I64) => {
545                                s.push_str("i32.reinterpret_f32\n");
546                                s.push_str("i64.extend_i32_u\n");
547                            }
548                            (CoreType::F64, CoreType::I64) => {
549                                s.push_str("i64.reinterpret_f64\n");
550                            }
551                            (a, b) if a == b => {}
552                            _ => unimplemented!("convert {from:?} to {to:?}"),
553                        }
554                        uwriteln!(s, "local.set $r{i}");
555                    }
556                }
557                s.push_str("br $r\n");
558                s.push_str("end\n");
559            }
560
561            // The catch-all block for out-of-bounds discriminants.
562            s.push_str("unreachable\n");
563            s.push_str("end\n");
564            for i in 0..lowered.len() {
565                uwriteln!(s, " local.get $r{i}");
566            }
567        };
568
569        match self {
570            Type::Bool
571            | Type::S8
572            | Type::U8
573            | Type::S16
574            | Type::U16
575            | Type::S32
576            | Type::U32
577            | Type::Char
578            | Type::S64
579            | Type::U64
580            | Type::Float32
581            | Type::Float64
582            | Type::String
583            | Type::List(_)
584            | Type::Map(_, _)
585            | Type::Flags(_)
586            | Type::Enum(_) => unreachable!(),
587
588            Type::Record(r) => record(s, helpers, r),
589            Type::Tuple(t) => record(s, helpers, t),
590            Type::Variant(v) => variant(
591                s,
592                helpers,
593                &v.iter().map(|t| t.as_ref()).collect::<Vec<_>>(),
594            ),
595            Type::Option(o) => variant(s, helpers, &[None, Some(&**o)]),
596            Type::Result { ok, err } => variant(s, helpers, &[ok.as_deref(), err.as_deref()]),
597        };
598        s.push_str(")\n");
599    }
600}
601
602#[derive(Clone)]
603enum FlatSource {
604    Local(u32),
605    LocalConvert {
606        local: u32,
607        from: CoreType,
608        to: CoreType,
609    },
610}
611
612impl fmt::Display for FlatSource {
613    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
614        match self {
615            FlatSource::Local(i) => write!(f, "(local.get {i})"),
616            FlatSource::LocalConvert { local, from, to } => {
617                match (from, to) {
618                    (a, b) if a == b => write!(f, "(local.get {local})"),
619                    (CoreType::I32, CoreType::F32) => {
620                        write!(f, "(f32.reinterpret_i32 (local.get {local}))")
621                    }
622                    (CoreType::I64, CoreType::I32) => {
623                        write!(f, "(i32.wrap_i64 (local.get {local}))")
624                    }
625                    (CoreType::I64, CoreType::F64) => {
626                        write!(f, "(f64.reinterpret_i64 (local.get {local}))")
627                    }
628                    (CoreType::I64, CoreType::F32) => {
629                        write!(
630                            f,
631                            "(f32.reinterpret_i32 (i32.wrap_i64 (local.get {local})))"
632                        )
633                    }
634                    _ => unimplemented!("convert {from:?} to {to:?}"),
635                }
636                // ..
637            }
638        }
639    }
640}
641
642fn lower_record<'a>(types: impl Iterator<Item = &'a Type>, vec: &mut Vec<CoreType>) {
643    for ty in types {
644        ty.lower(vec);
645    }
646}
647
648fn lower_variant<'a>(types: impl Iterator<Item = Option<&'a Type>>, vec: &mut Vec<CoreType>) {
649    vec.push(CoreType::I32);
650    let offset = vec.len();
651    for ty in types {
652        let ty = match ty {
653            Some(ty) => ty,
654            None => continue,
655        };
656        for (index, ty) in ty.lowered().iter().enumerate() {
657            let index = offset + index;
658            if index < vec.len() {
659                vec[index] = vec[index].join(*ty);
660            } else {
661                vec.push(*ty)
662            }
663        }
664    }
665}
666
667fn u32_count_from_flag_count(count: usize) -> usize {
668    match FlagsSize::from_count(count) {
669        FlagsSize::Size0 => 0,
670        FlagsSize::Size1 | FlagsSize::Size2 => 1,
671        FlagsSize::Size4Plus(n) => n.into(),
672    }
673}
674
675struct SizeAndAlignment {
676    size: usize,
677    alignment: u32,
678}
679
680impl Type {
681    fn lowered(&self) -> Vec<CoreType> {
682        let mut vec = Vec::new();
683        self.lower(&mut vec);
684        vec
685    }
686
687    fn lower(&self, vec: &mut Vec<CoreType>) {
688        match self {
689            Type::Bool
690            | Type::U8
691            | Type::S8
692            | Type::S16
693            | Type::U16
694            | Type::S32
695            | Type::U32
696            | Type::Char
697            | Type::Enum(_) => vec.push(CoreType::I32),
698            Type::S64 | Type::U64 => vec.push(CoreType::I64),
699            Type::Float32 => vec.push(CoreType::F32),
700            Type::Float64 => vec.push(CoreType::F64),
701            Type::String | Type::List(_) | Type::Map(_, _) => {
702                vec.push(CoreType::I32);
703                vec.push(CoreType::I32);
704            }
705            Type::Record(types) => lower_record(types.iter(), vec),
706            Type::Tuple(types) => lower_record(types.0.iter(), vec),
707            Type::Variant(types) => lower_variant(types.0.iter().map(|t| t.as_ref()), vec),
708            Type::Option(ty) => lower_variant([None, Some(&**ty)].into_iter(), vec),
709            Type::Result { ok, err } => {
710                lower_variant([ok.as_deref(), err.as_deref()].into_iter(), vec)
711            }
712            Type::Flags(count) => vec.extend(
713                iter::repeat(CoreType::I32).take(u32_count_from_flag_count(*count as usize)),
714            ),
715        }
716    }
717
718    fn size_and_alignment(&self) -> SizeAndAlignment {
719        match self {
720            Type::Bool | Type::S8 | Type::U8 => SizeAndAlignment {
721                size: 1,
722                alignment: 1,
723            },
724
725            Type::S16 | Type::U16 => SizeAndAlignment {
726                size: 2,
727                alignment: 2,
728            },
729
730            Type::S32 | Type::U32 | Type::Char | Type::Float32 => SizeAndAlignment {
731                size: 4,
732                alignment: 4,
733            },
734
735            Type::S64 | Type::U64 | Type::Float64 => SizeAndAlignment {
736                size: 8,
737                alignment: 8,
738            },
739
740            Type::String | Type::List(_) | Type::Map(_, _) => SizeAndAlignment {
741                size: 8,
742                alignment: 4,
743            },
744
745            Type::Record(types) => record_size_and_alignment(types.iter()),
746
747            Type::Tuple(types) => record_size_and_alignment(types.0.iter()),
748
749            Type::Variant(types) => variant_size_and_alignment(types.0.iter().map(|t| t.as_ref())),
750
751            Type::Enum(count) => variant_size_and_alignment((0..*count).map(|_| None)),
752
753            Type::Option(ty) => variant_size_and_alignment([None, Some(&**ty)].into_iter()),
754
755            Type::Result { ok, err } => {
756                variant_size_and_alignment([ok.as_deref(), err.as_deref()].into_iter())
757            }
758
759            Type::Flags(count) => match FlagsSize::from_count(*count as usize) {
760                FlagsSize::Size0 => SizeAndAlignment {
761                    size: 0,
762                    alignment: 1,
763                },
764                FlagsSize::Size1 => SizeAndAlignment {
765                    size: 1,
766                    alignment: 1,
767                },
768                FlagsSize::Size2 => SizeAndAlignment {
769                    size: 2,
770                    alignment: 2,
771                },
772                FlagsSize::Size4Plus(n) => SizeAndAlignment {
773                    size: usize::from(n) * 4,
774                    alignment: 4,
775                },
776            },
777        }
778    }
779}
780
781fn align_to(a: usize, align: u32) -> usize {
782    let align = align as usize;
783    (a + (align - 1)) & !(align - 1)
784}
785
786fn record_field_offsets<'a>(
787    types: impl IntoIterator<Item = &'a Type>,
788) -> impl Iterator<Item = (u32, &'a Type)> {
789    let mut offset = 0;
790    types.into_iter().map(move |ty| {
791        let SizeAndAlignment { size, alignment } = ty.size_and_alignment();
792        let ret = align_to(offset, alignment);
793        offset = ret + size;
794        (ret as u32, ty)
795    })
796}
797
798fn record_size_and_alignment<'a>(types: impl IntoIterator<Item = &'a Type>) -> SizeAndAlignment {
799    let mut offset = 0;
800    let mut align = 1;
801    for ty in types {
802        let SizeAndAlignment { size, alignment } = ty.size_and_alignment();
803        offset = align_to(offset, alignment) + size;
804        align = align.max(alignment);
805    }
806
807    SizeAndAlignment {
808        size: align_to(offset, align),
809        alignment: align,
810    }
811}
812
813fn variant_size_and_alignment<'a>(
814    types: impl ExactSizeIterator<Item = Option<&'a Type>>,
815) -> SizeAndAlignment {
816    let discriminant_size = DiscriminantSize::from_count(types.len()).unwrap();
817    let mut alignment = u32::from(discriminant_size);
818    let mut size = 0;
819    for ty in types {
820        if let Some(ty) = ty {
821            let size_and_alignment = ty.size_and_alignment();
822            alignment = alignment.max(size_and_alignment.alignment);
823            size = size.max(size_and_alignment.size);
824        }
825    }
826
827    SizeAndAlignment {
828        size: align_to(
829            align_to(usize::from(discriminant_size), alignment) + size,
830            alignment,
831        ),
832        alignment,
833    }
834}
835
836fn variant_memory_info<'a>(
837    types: impl ExactSizeIterator<Item = Option<&'a Type>>,
838) -> (DiscriminantSize, u32) {
839    let discriminant_size = DiscriminantSize::from_count(types.len()).unwrap();
840    let mut alignment = u32::from(discriminant_size);
841    for ty in types {
842        if let Some(ty) = ty {
843            let size_and_alignment = ty.size_and_alignment();
844            alignment = alignment.max(size_and_alignment.alignment);
845        }
846    }
847
848    (
849        discriminant_size,
850        align_to(usize::from(discriminant_size), alignment) as u32,
851    )
852}
853
854/// Generates the internals of a core wasm module which imports a single
855/// component function `IMPORT_FUNCTION` and exports a single component
856/// function `EXPORT_FUNCTION`.
857///
858/// The component function takes `params` as arguments and optionally returns
859/// `result`. The `lift_abi` and `lower_abi` fields indicate the ABI in-use for
860/// this operation.
861fn make_import_and_export(
862    params: &[&Type],
863    result: Option<&Type>,
864    lift_abi: LiftAbi,
865    lower_abi: LowerAbi,
866) -> String {
867    let params_lowered = params
868        .iter()
869        .flat_map(|ty| ty.lowered())
870        .collect::<Box<[_]>>();
871    let result_lowered = result.map(|t| t.lowered()).unwrap_or(Vec::new());
872
873    let mut wat = String::new();
874
875    enum Location {
876        Flat,
877        Indirect(u32),
878    }
879
880    // Generate the core wasm type corresponding to the imported function being
881    // lowered with `lower_abi`.
882    wat.push_str(&format!("(type $import (func"));
883    let max_import_params = match lower_abi {
884        LowerAbi::Sync => MAX_FLAT_PARAMS,
885        LowerAbi::Async => MAX_FLAT_ASYNC_PARAMS,
886    };
887    let (import_params_loc, nparams) = push_params(&mut wat, &params_lowered, max_import_params);
888    let import_results_loc = match lower_abi {
889        LowerAbi::Sync => {
890            push_result_or_retptr(&mut wat, &result_lowered, nparams, MAX_FLAT_RESULTS)
891        }
892        LowerAbi::Async => {
893            let loc = if result.is_none() {
894                Location::Flat
895            } else {
896                wat.push_str(" (param i32)"); // result pointer
897                Location::Indirect(nparams)
898            };
899            wat.push_str(" (result i32)"); // status code
900            loc
901        }
902    };
903    wat.push_str("))\n");
904
905    // Generate the import function.
906    wat.push_str(&format!(
907        r#"(import "host" "{IMPORT_FUNCTION}" (func $host (type $import)))"#
908    ));
909
910    // Do the same as above for the exported function's type which is lifted
911    // with `lift_abi`.
912    //
913    // Note that `export_results_loc` being `None` means that `task.return` is
914    // used to communicate results.
915    wat.push_str(&format!("(type $export (func"));
916    let (export_params_loc, _nparams) = push_params(&mut wat, &params_lowered, MAX_FLAT_PARAMS);
917    let export_results_loc = match lift_abi {
918        LiftAbi::Sync => Some(push_group(&mut wat, "result", &result_lowered, MAX_FLAT_RESULTS).0),
919        LiftAbi::AsyncCallback => {
920            wat.push_str(" (result i32)"); // status code
921            None
922        }
923        LiftAbi::AsyncStackful => None,
924    };
925    wat.push_str("))\n");
926
927    // If the export is async, generate `task.return` as an import as well
928    // which is necessary to communicate the results.
929    if export_results_loc.is_none() {
930        wat.push_str(&format!("(type $task.return (func"));
931        push_params(&mut wat, &result_lowered, MAX_FLAT_PARAMS);
932        wat.push_str("))\n");
933        wat.push_str(&format!(
934            r#"(import "" "task.return" (func $task.return (type $task.return)))"#
935        ));
936    }
937
938    wat.push_str(&format!(
939        r#"
940(func (export "{EXPORT_FUNCTION}") (type $export)
941    (local $retptr i32)
942    (local $argptr i32)
943        "#
944    ));
945    let mut store_helpers = IndexSet::new();
946    let mut load_helpers = IndexSet::new();
947
948    match (export_params_loc, import_params_loc) {
949        // flat => flat is just moving locals around
950        (Location::Flat, Location::Flat) => {
951            for (index, _) in params_lowered.iter().enumerate() {
952                uwrite!(wat, "local.get {index}\n");
953            }
954        }
955
956        // indirect => indirect is just moving locals around
957        (Location::Indirect(i), Location::Indirect(j)) => {
958            assert_eq!(j, 0);
959            uwrite!(wat, "local.get {i}\n");
960        }
961
962        // flat => indirect means that all parameters are stored in memory as
963        // if it was a record of all the parameters.
964        (Location::Flat, Location::Indirect(_)) => {
965            let SizeAndAlignment { size, alignment } =
966                record_size_and_alignment(params.iter().cloned());
967            wat.push_str(&format!(
968                r#"
969                    (local.set $argptr
970                        (call $realloc
971                            (i32.const 0)
972                            (i32.const 0)
973                            (i32.const {alignment})
974                            (i32.const {size})))
975                    local.get $argptr
976                "#
977            ));
978            let mut locals = (0..params_lowered.len() as u32).map(FlatSource::Local);
979            for (offset, ty) in record_field_offsets(params.iter().cloned()) {
980                ty.store_flat(&mut wat, "$argptr", offset, &mut locals, &mut store_helpers);
981            }
982            assert!(locals.next().is_none());
983        }
984
985        (Location::Indirect(_), Location::Flat) => unreachable!(),
986    }
987
988    // Pass a return-pointer if necessary.
989    match import_results_loc {
990        Location::Flat => {}
991        Location::Indirect(_) => {
992            let SizeAndAlignment { size, alignment } = result.unwrap().size_and_alignment();
993
994            wat.push_str(&format!(
995                r#"
996                    (local.set $retptr
997                        (call $realloc
998                            (i32.const 0)
999                            (i32.const 0)
1000                            (i32.const {alignment})
1001                            (i32.const {size})))
1002                    local.get $retptr
1003                "#
1004            ));
1005        }
1006    }
1007
1008    wat.push_str("call $host\n");
1009
1010    // Assert the lowered call is ready if an async code was returned.
1011    //
1012    // TODO: handle when the import isn't ready yet
1013    if let LowerAbi::Async = lower_abi {
1014        wat.push_str("i32.const 2\n");
1015        wat.push_str("i32.ne\n");
1016        wat.push_str("if unreachable end\n");
1017    }
1018
1019    // TODO: conditionally inject a yield here
1020
1021    match (import_results_loc, export_results_loc) {
1022        // flat => flat results involves nothing, the results are already on
1023        // the stack.
1024        (Location::Flat, Some(Location::Flat)) => {}
1025
1026        // indirect => indirect results requires returning the `$retptr` the
1027        // host call filled in.
1028        (Location::Indirect(_), Some(Location::Indirect(_))) => {
1029            wat.push_str("local.get $retptr\n");
1030        }
1031
1032        // indirect => flat requires loading the result from the return pointer
1033        (Location::Indirect(_), Some(Location::Flat)) => {
1034            result
1035                .unwrap()
1036                .load_flat(&mut wat, "$retptr", 0, &mut load_helpers);
1037        }
1038
1039        // flat => task.return is easy, the results are already there so just
1040        // call the function.
1041        (Location::Flat, None) => {
1042            wat.push_str("call $task.return\n");
1043        }
1044
1045        // indirect => task.return needs to forward `$retptr` if the results
1046        // are indirect, or otherwise it must be loaded from memory to a flat
1047        // representation.
1048        (Location::Indirect(_), None) => {
1049            if result_lowered.len() <= MAX_FLAT_PARAMS {
1050                result
1051                    .unwrap()
1052                    .load_flat(&mut wat, "$retptr", 0, &mut load_helpers);
1053            } else {
1054                wat.push_str("local.get $retptr\n");
1055            }
1056            wat.push_str("call $task.return\n");
1057        }
1058
1059        (Location::Flat, Some(Location::Indirect(_))) => unreachable!(),
1060    }
1061
1062    if let LiftAbi::AsyncCallback = lift_abi {
1063        wat.push_str("i32.const 0\n"); // completed status code
1064    }
1065
1066    wat.push_str(")\n");
1067
1068    // Generate a `callback` function for the callback ABI.
1069    //
1070    // TODO: fill this in
1071    if let LiftAbi::AsyncCallback = lift_abi {
1072        wat.push_str(
1073            r#"
1074(func (export "callback") (param i32 i32 i32) (result i32) unreachable)
1075            "#,
1076        );
1077    }
1078
1079    // Fill out all store/load helpers that were needed during generation
1080    // above. This is a fix-point-loop since each helper may end up requiring
1081    // more helpers.
1082    let mut i = 0;
1083    while i < store_helpers.len() {
1084        let ty = store_helpers[i].0;
1085        ty.store_flat_helper(&mut wat, i, &mut store_helpers);
1086        i += 1;
1087    }
1088    i = 0;
1089    while i < load_helpers.len() {
1090        let ty = load_helpers[i].0;
1091        ty.load_flat_helper(&mut wat, i, &mut load_helpers);
1092        i += 1;
1093    }
1094
1095    return wat;
1096
1097    fn push_params(wat: &mut String, params: &[CoreType], max_flat: usize) -> (Location, u32) {
1098        push_group(wat, "param", params, max_flat)
1099    }
1100
1101    fn push_group(
1102        wat: &mut String,
1103        name: &str,
1104        params: &[CoreType],
1105        max_flat: usize,
1106    ) -> (Location, u32) {
1107        let mut nparams = 0;
1108        let loc = if params.is_empty() {
1109            // nothing to emit...
1110            Location::Flat
1111        } else if params.len() <= max_flat {
1112            wat.push_str(&format!(" ({name}"));
1113            for ty in params {
1114                wat.push_str(&format!(" {ty}"));
1115                nparams += 1;
1116            }
1117            wat.push_str(")");
1118            Location::Flat
1119        } else {
1120            wat.push_str(&format!(" ({name} i32)"));
1121            nparams += 1;
1122            Location::Indirect(0)
1123        };
1124        (loc, nparams)
1125    }
1126
1127    fn push_result_or_retptr(
1128        wat: &mut String,
1129        results: &[CoreType],
1130        nparams: u32,
1131        max_flat: usize,
1132    ) -> Location {
1133        if results.is_empty() {
1134            // nothing to emit...
1135            Location::Flat
1136        } else if results.len() <= max_flat {
1137            wat.push_str(" (result");
1138            for ty in results {
1139                wat.push_str(&format!(" {ty}"));
1140            }
1141            wat.push_str(")");
1142            Location::Flat
1143        } else {
1144            wat.push_str(" (param i32)");
1145            Location::Indirect(nparams)
1146        }
1147    }
1148}
1149
1150struct Helper<'a>(&'a Type);
1151
1152impl Hash for Helper<'_> {
1153    fn hash<H: Hasher>(&self, h: &mut H) {
1154        std::ptr::hash(self.0, h);
1155    }
1156}
1157
1158impl PartialEq for Helper<'_> {
1159    fn eq(&self, other: &Self) -> bool {
1160        std::ptr::eq(self.0, other.0)
1161    }
1162}
1163
1164impl Eq for Helper<'_> {}
1165
1166fn make_rust_name(name_counter: &mut u32) -> Ident {
1167    let name = format_ident!("Foo{name_counter}");
1168    *name_counter += 1;
1169    name
1170}
1171
1172/// Generate a [`TokenStream`] containing the rust type name for a type.
1173///
1174/// The `name_counter` parameter is used to generate names for each recursively visited type.  The `declarations`
1175/// parameter is used to accumulate declarations for each recursively visited type.
1176pub fn rust_type(ty: &Type, name_counter: &mut u32, declarations: &mut TokenStream) -> TokenStream {
1177    match ty {
1178        Type::Bool => quote!(bool),
1179        Type::S8 => quote!(i8),
1180        Type::U8 => quote!(u8),
1181        Type::S16 => quote!(i16),
1182        Type::U16 => quote!(u16),
1183        Type::S32 => quote!(i32),
1184        Type::U32 => quote!(u32),
1185        Type::S64 => quote!(i64),
1186        Type::U64 => quote!(u64),
1187        Type::Float32 => quote!(Float32),
1188        Type::Float64 => quote!(Float64),
1189        Type::Char => quote!(char),
1190        Type::String => quote!(Box<str>),
1191        Type::List(ty) => {
1192            let ty = rust_type(ty, name_counter, declarations);
1193            quote!(Vec<#ty>)
1194        }
1195        Type::Map(key_ty, value_ty) => {
1196            let key_ty = rust_type(key_ty, name_counter, declarations);
1197            let value_ty = rust_type(value_ty, name_counter, declarations);
1198            quote!(std::collections::HashMap<#key_ty, #value_ty>)
1199        }
1200        Type::Record(types) => {
1201            let fields = types
1202                .iter()
1203                .enumerate()
1204                .map(|(index, ty)| {
1205                    let name = format_ident!("f{index}");
1206                    let ty = rust_type(ty, name_counter, declarations);
1207                    quote!(#name: #ty,)
1208                })
1209                .collect::<TokenStream>();
1210
1211            let name = make_rust_name(name_counter);
1212
1213            declarations.extend(quote! {
1214                #[derive(ComponentType, Lift, Lower, PartialEq, Debug, Clone, Arbitrary)]
1215                #[component(record)]
1216                struct #name {
1217                    #fields
1218                }
1219            });
1220
1221            quote!(#name)
1222        }
1223        Type::Tuple(types) => {
1224            let fields = types
1225                .0
1226                .iter()
1227                .map(|ty| {
1228                    let ty = rust_type(ty, name_counter, declarations);
1229                    quote!(#ty,)
1230                })
1231                .collect::<TokenStream>();
1232
1233            quote!((#fields))
1234        }
1235        Type::Variant(types) => {
1236            let cases = types
1237                .0
1238                .iter()
1239                .enumerate()
1240                .map(|(index, ty)| {
1241                    let name = format_ident!("C{index}");
1242                    let ty = match ty {
1243                        Some(ty) => {
1244                            let ty = rust_type(ty, name_counter, declarations);
1245                            quote!((#ty))
1246                        }
1247                        None => quote!(),
1248                    };
1249                    quote!(#name #ty,)
1250                })
1251                .collect::<TokenStream>();
1252
1253            let name = make_rust_name(name_counter);
1254            declarations.extend(quote! {
1255                #[derive(ComponentType, Lift, Lower, PartialEq, Debug, Clone, Arbitrary)]
1256                #[component(variant)]
1257                enum #name {
1258                    #cases
1259                }
1260            });
1261
1262            quote!(#name)
1263        }
1264        Type::Enum(count) => {
1265            let cases = (0..*count)
1266                .map(|index| {
1267                    let name = format_ident!("E{index}");
1268                    quote!(#name,)
1269                })
1270                .collect::<TokenStream>();
1271
1272            let name = make_rust_name(name_counter);
1273            let repr = match DiscriminantSize::from_count(*count as usize).unwrap() {
1274                DiscriminantSize::Size1 => quote!(u8),
1275                DiscriminantSize::Size2 => quote!(u16),
1276                DiscriminantSize::Size4 => quote!(u32),
1277            };
1278
1279            declarations.extend(quote! {
1280                #[derive(ComponentType, Lift, Lower, PartialEq, Eq, Hash, Debug, Copy, Clone, Arbitrary)]
1281                #[component(enum)]
1282                #[repr(#repr)]
1283                enum #name {
1284                    #cases
1285                }
1286            });
1287
1288            quote!(#name)
1289        }
1290        Type::Option(ty) => {
1291            let ty = rust_type(ty, name_counter, declarations);
1292            quote!(Option<#ty>)
1293        }
1294        Type::Result { ok, err } => {
1295            let ok = match ok {
1296                Some(ok) => rust_type(ok, name_counter, declarations),
1297                None => quote!(()),
1298            };
1299            let err = match err {
1300                Some(err) => rust_type(err, name_counter, declarations),
1301                None => quote!(()),
1302            };
1303            quote!(Result<#ok, #err>)
1304        }
1305        Type::Flags(count) => {
1306            let type_name = make_rust_name(name_counter);
1307
1308            let mut flags = TokenStream::new();
1309            let mut names = TokenStream::new();
1310
1311            for index in 0..*count {
1312                let name = format_ident!("F{index}");
1313                flags.extend(quote!(const #name;));
1314                names.extend(quote!(#type_name::#name,))
1315            }
1316
1317            declarations.extend(quote! {
1318                wasmtime::component::flags! {
1319                    #type_name {
1320                        #flags
1321                    }
1322                }
1323
1324                impl<'a> arbitrary::Arbitrary<'a> for #type_name {
1325                    fn arbitrary(input: &mut arbitrary::Unstructured<'a>) -> arbitrary::Result<Self> {
1326                        let mut flags = #type_name::default();
1327                        for flag in [#names] {
1328                            if input.arbitrary()? {
1329                                flags |= flag;
1330                            }
1331                        }
1332                        Ok(flags)
1333                    }
1334                }
1335            });
1336
1337            quote!(#type_name)
1338        }
1339    }
1340}
1341
1342#[derive(Default)]
1343struct TypesBuilder<'a> {
1344    next: u32,
1345    worklist: Vec<(u32, &'a Type)>,
1346}
1347
1348impl<'a> TypesBuilder<'a> {
1349    fn write_ref(&mut self, ty: &'a Type, dst: &mut String) {
1350        match ty {
1351            // Primitive types can be referenced directly
1352            Type::Bool => dst.push_str("bool"),
1353            Type::S8 => dst.push_str("s8"),
1354            Type::U8 => dst.push_str("u8"),
1355            Type::S16 => dst.push_str("s16"),
1356            Type::U16 => dst.push_str("u16"),
1357            Type::S32 => dst.push_str("s32"),
1358            Type::U32 => dst.push_str("u32"),
1359            Type::S64 => dst.push_str("s64"),
1360            Type::U64 => dst.push_str("u64"),
1361            Type::Float32 => dst.push_str("float32"),
1362            Type::Float64 => dst.push_str("float64"),
1363            Type::Char => dst.push_str("char"),
1364            Type::String => dst.push_str("string"),
1365
1366            // Otherwise emit a reference to the type and remember to generate
1367            // the corresponding type alias later.
1368            Type::List(_)
1369            | Type::Map(_, _)
1370            | Type::Record(_)
1371            | Type::Tuple(_)
1372            | Type::Variant(_)
1373            | Type::Enum(_)
1374            | Type::Option(_)
1375            | Type::Result { .. }
1376            | Type::Flags(_) => {
1377                let idx = self.next;
1378                self.next += 1;
1379                uwrite!(dst, "$t{idx}");
1380                self.worklist.push((idx, ty));
1381            }
1382        }
1383    }
1384
1385    fn write_decl(&mut self, idx: u32, ty: &'a Type) -> String {
1386        let mut decl = format!("(type $t{idx}' ");
1387        match ty {
1388            Type::Bool
1389            | Type::S8
1390            | Type::U8
1391            | Type::S16
1392            | Type::U16
1393            | Type::S32
1394            | Type::U32
1395            | Type::S64
1396            | Type::U64
1397            | Type::Float32
1398            | Type::Float64
1399            | Type::Char
1400            | Type::String => unreachable!(),
1401
1402            Type::List(ty) => {
1403                decl.push_str("(list ");
1404                self.write_ref(ty, &mut decl);
1405                decl.push_str(")");
1406            }
1407            Type::Map(key_ty, value_ty) => {
1408                decl.push_str("(map ");
1409                self.write_ref(key_ty, &mut decl);
1410                decl.push_str(" ");
1411                self.write_ref(value_ty, &mut decl);
1412                decl.push_str(")");
1413            }
1414            Type::Record(types) => {
1415                decl.push_str("(record");
1416                for (index, ty) in types.iter().enumerate() {
1417                    uwrite!(decl, r#" (field "f{index}" "#);
1418                    self.write_ref(ty, &mut decl);
1419                    decl.push_str(")");
1420                }
1421                decl.push_str(")");
1422            }
1423            Type::Tuple(types) => {
1424                decl.push_str("(tuple");
1425                for ty in types.iter() {
1426                    decl.push_str(" ");
1427                    self.write_ref(ty, &mut decl);
1428                }
1429                decl.push_str(")");
1430            }
1431            Type::Variant(types) => {
1432                decl.push_str("(variant");
1433                for (index, ty) in types.iter().enumerate() {
1434                    uwrite!(decl, r#" (case "C{index}""#);
1435                    if let Some(ty) = ty {
1436                        decl.push_str(" ");
1437                        self.write_ref(ty, &mut decl);
1438                    }
1439                    decl.push_str(")");
1440                }
1441                decl.push_str(")");
1442            }
1443            Type::Enum(count) => {
1444                decl.push_str("(enum");
1445                for index in 0..*count {
1446                    uwrite!(decl, r#" "E{index}""#);
1447                }
1448                decl.push_str(")");
1449            }
1450            Type::Option(ty) => {
1451                decl.push_str("(option ");
1452                self.write_ref(ty, &mut decl);
1453                decl.push_str(")");
1454            }
1455            Type::Result { ok, err } => {
1456                decl.push_str("(result");
1457                if let Some(ok) = ok {
1458                    decl.push_str(" ");
1459                    self.write_ref(ok, &mut decl);
1460                }
1461                if let Some(err) = err {
1462                    decl.push_str(" (error ");
1463                    self.write_ref(err, &mut decl);
1464                    decl.push_str(")");
1465                }
1466                decl.push_str(")");
1467            }
1468            Type::Flags(count) => {
1469                decl.push_str("(flags");
1470                for index in 0..*count {
1471                    uwrite!(decl, r#" "F{index}""#);
1472                }
1473                decl.push_str(")");
1474            }
1475        }
1476        decl.push_str(")\n");
1477        uwriteln!(decl, "(import \"t{idx}\" (type $t{idx} (eq $t{idx}')))");
1478        decl
1479    }
1480}
1481
1482/// Represents custom fragments of a WAT file which may be used to create a component for exercising [`TestCase`]s
1483#[derive(Debug)]
1484pub struct Declarations {
1485    /// Type declarations (if any) referenced by `params` and/or `result`
1486    pub types: Cow<'static, str>,
1487    /// Types to thread through when instantiating sub-components.
1488    pub type_instantiation_args: Cow<'static, str>,
1489    /// Parameter declarations used for the imported and exported functions
1490    pub params: Cow<'static, str>,
1491    /// Result declaration used for the imported and exported functions
1492    pub results: Cow<'static, str>,
1493    /// Implementation of the "caller" component, which invokes the `callee`
1494    /// composed component.
1495    pub caller_module: Cow<'static, str>,
1496    /// Implementation of the "callee" component, which invokes the host.
1497    pub callee_module: Cow<'static, str>,
1498    /// Options used for caller/calle ABI/etc.
1499    pub options: TestCaseOptions,
1500}
1501
1502impl Declarations {
1503    /// Generate a complete WAT file based on the specified fragments.
1504    pub fn make_component(&self) -> Box<str> {
1505        let Self {
1506            types,
1507            type_instantiation_args,
1508            params,
1509            results,
1510            caller_module,
1511            callee_module,
1512            options,
1513        } = self;
1514        let mk_component = |name: &str,
1515                            module: &str,
1516                            import_async: bool,
1517                            export_async: bool,
1518                            encoding: StringEncoding,
1519                            lift_abi: LiftAbi,
1520                            lower_abi: LowerAbi| {
1521            let import_async = if import_async { "async" } else { "" };
1522            let export_async = if export_async { "async" } else { "" };
1523            let lower_async_option = match lower_abi {
1524                LowerAbi::Sync => "",
1525                LowerAbi::Async => "async",
1526            };
1527            let lift_async_option = match lift_abi {
1528                LiftAbi::Sync => "",
1529                LiftAbi::AsyncStackful => "async",
1530                LiftAbi::AsyncCallback => "async (callback (core func $i \"callback\"))",
1531            };
1532
1533            let mut intrinsic_defs = String::new();
1534            let mut intrinsic_imports = String::new();
1535
1536            match lift_abi {
1537                LiftAbi::Sync => {}
1538                LiftAbi::AsyncCallback | LiftAbi::AsyncStackful => {
1539                    intrinsic_defs.push_str(&format!(
1540                        r#"
1541(core func $task.return (canon task.return {results}
1542    (memory (core memory $libc "memory")) string-encoding={encoding}))
1543                        "#,
1544                    ));
1545                    intrinsic_imports.push_str(
1546                        r#"
1547(with "" (instance (export "task.return" (func $task.return))))
1548                        "#,
1549                    );
1550                }
1551            }
1552
1553            format!(
1554                r#"
1555(component ${name}
1556    {types}
1557    (type $import_sig (func {import_async} {params} {results}))
1558    (type $export_sig (func {export_async} {params} {results}))
1559    (import "{IMPORT_FUNCTION}" (func $f (type $import_sig)))
1560
1561    (core instance $libc (instantiate $libc))
1562
1563    (core func $f_lower (canon lower
1564        (func $f)
1565        (memory (core memory $libc "memory"))
1566        (realloc (core func $libc "realloc"))
1567        string-encoding={encoding}
1568        {lower_async_option}
1569    ))
1570
1571    {intrinsic_defs}
1572
1573    (core module $m
1574        (memory (import "libc" "memory") 1)
1575        (func $realloc (import "libc" "realloc") (param i32 i32 i32 i32) (result i32))
1576
1577        {module}
1578    )
1579
1580    (core instance $i (instantiate $m
1581        (with "libc" (instance $libc))
1582        (with "host" (instance (export "{IMPORT_FUNCTION}" (func $f_lower))))
1583        {intrinsic_imports}
1584    ))
1585
1586    (func (export "{EXPORT_FUNCTION}") (type $export_sig)
1587        (canon lift
1588            (core func $i "{EXPORT_FUNCTION}")
1589            (memory (core memory $libc "memory"))
1590            (realloc (core func $libc "realloc"))
1591            string-encoding={encoding}
1592            {lift_async_option}
1593        )
1594    )
1595)
1596            "#
1597            )
1598        };
1599
1600        let c1 = mk_component(
1601            "callee",
1602            &callee_module,
1603            options.host_async,
1604            options.guest_callee_async,
1605            options.callee_encoding,
1606            options.callee_lift_abi,
1607            options.callee_lower_abi,
1608        );
1609        let c2 = mk_component(
1610            "caller",
1611            &caller_module,
1612            options.guest_callee_async,
1613            options.guest_caller_async,
1614            options.caller_encoding,
1615            options.caller_lift_abi,
1616            options.caller_lower_abi,
1617        );
1618        let host_async = if options.host_async { "async" } else { "" };
1619
1620        format!(
1621            r#"
1622            (component
1623                (core module $libc
1624                    (memory (export "memory") 1)
1625                    {REALLOC_AND_FREE}
1626                )
1627
1628
1629                {types}
1630
1631                (type $host_sig (func {host_async} {params} {results}))
1632                (import "{IMPORT_FUNCTION}" (func $f (type $host_sig)))
1633
1634                {c1}
1635                {c2}
1636                (instance $c1 (instantiate $callee
1637                    {type_instantiation_args}
1638                    (with "{IMPORT_FUNCTION}" (func $f))
1639                ))
1640                (instance $c2 (instantiate $caller
1641                    {type_instantiation_args}
1642                    (with "{IMPORT_FUNCTION}" (func $c1 "{EXPORT_FUNCTION}"))
1643                ))
1644                (export "{EXPORT_FUNCTION}" (func $c2 "{EXPORT_FUNCTION}"))
1645            )"#,
1646        )
1647        .into()
1648    }
1649}
1650
1651/// Represents a test case for calling a component function
1652#[derive(Debug)]
1653pub struct TestCase<'a> {
1654    /// The types of parameters to pass to the function
1655    pub params: Vec<&'a Type>,
1656    /// The result types of the function
1657    pub result: Option<&'a Type>,
1658    /// ABI options to use for this test case.
1659    pub options: TestCaseOptions,
1660}
1661
1662/// Collection of options which configure how the caller/callee/etc ABIs are
1663/// all configured.
1664#[derive(Debug, Arbitrary, Copy, Clone)]
1665pub struct TestCaseOptions {
1666    /// Whether or not the guest caller component (the entrypoint) is using an
1667    /// `async` function type.
1668    pub guest_caller_async: bool,
1669    /// Whether or not the guest callee component (what the entrypoint calls)
1670    /// is using an `async` function type.
1671    pub guest_callee_async: bool,
1672    /// Whether or not the host is using an async function type (what the
1673    /// guest callee calls).
1674    pub host_async: bool,
1675    /// The string encoding of the caller component.
1676    pub caller_encoding: StringEncoding,
1677    /// The string encoding of the callee component.
1678    pub callee_encoding: StringEncoding,
1679    /// The ABI that the caller component is using to lift its export (the main
1680    /// entrypoint).
1681    pub caller_lift_abi: LiftAbi,
1682    /// The ABI that the callee component is using to lift its export (called
1683    /// by the caller).
1684    pub callee_lift_abi: LiftAbi,
1685    /// The ABI that the caller component is using to lower its import (the
1686    /// callee's export).
1687    pub caller_lower_abi: LowerAbi,
1688    /// The ABI that the callee component is using to lower its import (the
1689    /// host function).
1690    pub callee_lower_abi: LowerAbi,
1691}
1692
1693#[derive(Debug, Arbitrary, Copy, Clone)]
1694pub enum LiftAbi {
1695    Sync,
1696    AsyncStackful,
1697    AsyncCallback,
1698}
1699
1700#[derive(Debug, Arbitrary, Copy, Clone)]
1701pub enum LowerAbi {
1702    Sync,
1703    Async,
1704}
1705
1706impl LiftAbi {
1707    fn is_async(self) -> bool {
1708        !matches!(self, Self::Sync)
1709    }
1710
1711    fn ensure_async(&mut self) {
1712        if !self.is_async() {
1713            *self = Self::AsyncStackful;
1714        }
1715    }
1716}
1717
1718impl LowerAbi {
1719    fn is_async(self) -> bool {
1720        matches!(self, Self::Async)
1721    }
1722
1723    fn ensure_async(&mut self) {
1724        *self = Self::Async;
1725    }
1726}
1727
1728impl<'a> TestCase<'a> {
1729    pub fn generate(types: &'a [Type], u: &mut Unstructured<'_>) -> arbitrary::Result<Self> {
1730        let max_params = if types.len() > 0 { 5 } else { 0 };
1731        let params = (0..u.int_in_range(0..=max_params)?)
1732            .map(|_| u.choose(&types))
1733            .collect::<arbitrary::Result<Vec<_>>>()?;
1734        let result = if types.len() > 0 && u.arbitrary()? {
1735            Some(u.choose(&types)?)
1736        } else {
1737            None
1738        };
1739
1740        let mut options = u.arbitrary::<TestCaseOptions>()?;
1741
1742        // Keep each boundary self-consistent: async function types require async
1743        // lowering/lifting, and async callees force their callers async too.
1744        if options.host_async || options.callee_lower_abi.is_async() {
1745            options.host_async = true;
1746            options.callee_lower_abi.ensure_async();
1747            options.guest_callee_async = true;
1748        }
1749        if options.guest_callee_async
1750            || options.callee_lift_abi.is_async()
1751            || options.caller_lower_abi.is_async()
1752        {
1753            options.guest_callee_async = true;
1754            options.callee_lift_abi.ensure_async();
1755            options.caller_lower_abi.ensure_async();
1756            options.guest_caller_async = true;
1757        }
1758        if options.guest_caller_async || options.caller_lift_abi.is_async() {
1759            options.guest_caller_async = true;
1760            options.caller_lift_abi.ensure_async();
1761        }
1762
1763        Ok(Self {
1764            params,
1765            result,
1766            options,
1767        })
1768    }
1769
1770    /// Generate a `Declarations` for this `TestCase` which may be used to build a component to execute the case.
1771    pub fn declarations(&self) -> Declarations {
1772        let mut builder = TypesBuilder::default();
1773
1774        let mut params = String::new();
1775        for (i, ty) in self.params.iter().enumerate() {
1776            params.push_str(&format!(" (param \"p{i}\" "));
1777            builder.write_ref(ty, &mut params);
1778            params.push_str(")");
1779        }
1780
1781        let mut results = String::new();
1782        if let Some(ty) = self.result {
1783            results.push_str(&format!(" (result "));
1784            builder.write_ref(ty, &mut results);
1785            results.push_str(")");
1786        }
1787
1788        let caller_module = make_import_and_export(
1789            &self.params,
1790            self.result,
1791            self.options.caller_lift_abi,
1792            self.options.caller_lower_abi,
1793        );
1794        let callee_module = make_import_and_export(
1795            &self.params,
1796            self.result,
1797            self.options.callee_lift_abi,
1798            self.options.callee_lower_abi,
1799        );
1800
1801        let mut type_decls = Vec::new();
1802        let mut type_instantiation_args = String::new();
1803        while let Some((idx, ty)) = builder.worklist.pop() {
1804            type_decls.push(builder.write_decl(idx, ty));
1805            uwriteln!(type_instantiation_args, "(with \"t{idx}\" (type $t{idx}))");
1806        }
1807
1808        // Note that types are printed here in reverse order since they were
1809        // pushed onto `type_decls` as they were referenced meaning the last one
1810        // is the "base" one.
1811        let mut types = String::new();
1812        for decl in type_decls.into_iter().rev() {
1813            types.push_str(&decl);
1814            types.push_str("\n");
1815        }
1816
1817        Declarations {
1818            types: types.into(),
1819            type_instantiation_args: type_instantiation_args.into(),
1820            params: params.into(),
1821            results: results.into(),
1822            caller_module: caller_module.into(),
1823            callee_module: callee_module.into(),
1824            options: self.options,
1825        }
1826    }
1827}
1828
1829#[derive(Copy, Clone, Debug, Arbitrary)]
1830pub enum StringEncoding {
1831    Utf8,
1832    Utf16,
1833    Latin1OrUtf16,
1834}
1835
1836impl fmt::Display for StringEncoding {
1837    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1838        match self {
1839            StringEncoding::Utf8 => fmt::Display::fmt(&"utf8", f),
1840            StringEncoding::Utf16 => fmt::Display::fmt(&"utf16", f),
1841            StringEncoding::Latin1OrUtf16 => fmt::Display::fmt(&"latin1+utf16", f),
1842        }
1843    }
1844}
1845
1846impl ToTokens for TestCaseOptions {
1847    fn to_tokens(&self, tokens: &mut TokenStream) {
1848        let TestCaseOptions {
1849            guest_caller_async,
1850            guest_callee_async,
1851            host_async,
1852            caller_encoding,
1853            callee_encoding,
1854            caller_lift_abi,
1855            callee_lift_abi,
1856            caller_lower_abi,
1857            callee_lower_abi,
1858        } = self;
1859        tokens.extend(quote!(wasmtime_test_util::component_fuzz::TestCaseOptions {
1860            guest_caller_async: #guest_caller_async,
1861            guest_callee_async: #guest_callee_async,
1862            host_async: #host_async,
1863            caller_encoding: #caller_encoding,
1864            callee_encoding: #callee_encoding,
1865            caller_lift_abi: #caller_lift_abi,
1866            callee_lift_abi: #callee_lift_abi,
1867            caller_lower_abi: #caller_lower_abi,
1868            callee_lower_abi: #callee_lower_abi,
1869        }));
1870    }
1871}
1872
1873impl ToTokens for LowerAbi {
1874    fn to_tokens(&self, tokens: &mut TokenStream) {
1875        let me = match self {
1876            LowerAbi::Sync => quote!(Sync),
1877            LowerAbi::Async => quote!(Async),
1878        };
1879        tokens.extend(quote!(wasmtime_test_util::component_fuzz::LowerAbi::#me));
1880    }
1881}
1882
1883impl ToTokens for LiftAbi {
1884    fn to_tokens(&self, tokens: &mut TokenStream) {
1885        let me = match self {
1886            LiftAbi::Sync => quote!(Sync),
1887            LiftAbi::AsyncCallback => quote!(AsyncCallback),
1888            LiftAbi::AsyncStackful => quote!(AsyncStackful),
1889        };
1890        tokens.extend(quote!(wasmtime_test_util::component_fuzz::LiftAbi::#me));
1891    }
1892}
1893
1894impl ToTokens for StringEncoding {
1895    fn to_tokens(&self, tokens: &mut TokenStream) {
1896        let me = match self {
1897            StringEncoding::Utf8 => quote!(Utf8),
1898            StringEncoding::Utf16 => quote!(Utf16),
1899            StringEncoding::Latin1OrUtf16 => quote!(Latin1OrUtf16),
1900        };
1901        tokens.extend(quote!(wasmtime_test_util::component_fuzz::StringEncoding::#me));
1902    }
1903}
1904
1905#[cfg(test)]
1906mod tests {
1907    use super::*;
1908
1909    #[test]
1910    fn arbtest() {
1911        arbtest::arbtest(|u| {
1912            let mut fuel = 100;
1913            let types = (0..5)
1914                .map(|_| Type::generate(u, 3, &mut fuel))
1915                .collect::<arbitrary::Result<Vec<_>>>()?;
1916            let case = TestCase::generate(&types, u)?;
1917            let decls = case.declarations();
1918            let component = decls.make_component();
1919            let wasm = wat::parse_str(&component).unwrap_or_else(|e| {
1920                panic!("failed to parse generated component as wat: {e}\n\n{component}");
1921            });
1922            wasmparser::Validator::new_with_features(wasmparser::WasmFeatures::all())
1923                .validate_all(&wasm)
1924                .unwrap_or_else(|e| {
1925                    let mut wat = String::new();
1926                    let mut dst = wasmprinter::PrintFmtWrite(&mut wat);
1927                    let to_print = if wasmprinter::Config::new()
1928                        .print_offsets(true)
1929                        .print_operand_stack(true)
1930                        .print(&wasm, &mut dst)
1931                        .is_ok()
1932                    {
1933                        &wat[..]
1934                    } else {
1935                        &component[..]
1936                    };
1937                    panic!("generated component is not valid wasm: {e}\n\n{to_print}");
1938                });
1939            Ok(())
1940        })
1941        .budget_ms(1_000)
1942        // .seed(0x3c9050d4000000e9)
1943        ;
1944    }
1945}