From 20b36e42e016bdf47e6b12f9c11c479534127709 Mon Sep 17 00:00:00 2001 From: Martin Evans Date: Sun, 17 Aug 2025 21:41:26 +0100 Subject: [PATCH 1/9] Initial sketching out what's required for wasm component support --- src/Components/Component.cs | 189 ++++++++++++++++++++++++++++ src/Components/ComponentExport.cs | 57 +++++++++ src/Components/ComponentFunction.cs | 19 +++ src/Components/ComponentInstance.cs | 16 +++ src/Components/ComponentLinker.cs | 75 +++++++++++ src/Components/ComponentValue.cs | 130 +++++++++++++++++++ src/Module.cs | 33 +++-- src/Wasmtime.csproj | 21 +--- 8 files changed, 510 insertions(+), 30 deletions(-) create mode 100644 src/Components/Component.cs create mode 100644 src/Components/ComponentExport.cs create mode 100644 src/Components/ComponentFunction.cs create mode 100644 src/Components/ComponentInstance.cs create mode 100644 src/Components/ComponentLinker.cs create mode 100644 src/Components/ComponentValue.cs diff --git a/src/Components/Component.cs b/src/Components/Component.cs new file mode 100644 index 00000000..1dfcd573 --- /dev/null +++ b/src/Components/Component.cs @@ -0,0 +1,189 @@ +using Microsoft.Win32.SafeHandles; +using System; +using System.Runtime.InteropServices; + +namespace Wasmtime.Components; + +/// +/// Representation of a component in the component model. +/// +public class Component + : IDisposable +{ + private readonly Handle handle; + + internal Handle NativeHandle + { + get + { + if (handle.IsInvalid || handle.IsClosed) + { + throw new ObjectDisposedException(typeof(Module).FullName); + } + + return handle; + } + } + + internal Component(IntPtr handle) + { + this.handle = new Handle(handle); + } + + /// + public void Dispose() + { + handle.Dispose(); + } + + /// + /// Creates a given bytes. + /// + /// The engine to use for the Component. + /// The bytes of the Component. + /// Returns a new . + public static Component FromBytes(Engine engine, ReadOnlySpan bytes) + { + if (engine is null) + { + throw new ArgumentNullException(nameof(engine)); + } + + unsafe + { + fixed (byte* ptr = bytes) + { + var error = Native.wasmtime_component_new(engine.NativeHandle, ptr, (UIntPtr)bytes.Length, out var handle); + if (error != IntPtr.Zero) + { + throw new WasmtimeException($"WebAssembly component is not valid: {WasmtimeException.FromOwnedError(error).Message}"); + } + + return new Component(handle); + } + } + } + + /// + /// This function serializes compiled component artifacts as blob data. + /// + /// If the conversion is successful, the serialized compiled component. + public byte[] Serialize() + { + var error = Native.wasmtime_component_serialize(NativeHandle, out var bytes); + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + using (bytes) + return bytes.ToArray(); + } + + /// + /// Deserializes a previously serialized component from a span of bytes. + /// + /// The engine to use to deserialize the component. + /// The previously serialized component bytes. + /// Returns the that was previously serialized. + /// The passed bytes must come from a previous call to . + public static Component Deserialize(Engine engine, ReadOnlySpan bytes) + { + if (engine is null) + { + throw new ArgumentNullException(nameof(engine)); + } + + unsafe + { + fixed (byte* ptr = bytes) + { + var error = Native.wasmtime_component_deserialize(engine.NativeHandle, ptr, (UIntPtr)bytes.Length, out var handle); + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + return new Component(handle); + } + } + } + + /// + /// Deserializes a previously serialized component from a file. + /// + /// The engine to deserialize the component with. + /// The path to the previously serialized component. + /// Returns the that was previously serialized. + /// The file's contents must come from a previous call to . + public static Component DeserializeFile(Engine engine, string path) + { + if (engine is null) + { + throw new ArgumentNullException(nameof(engine)); + } + + var error = Native.wasmtime_component_deserialize_file(engine.NativeHandle, path, out var handle); + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + return new Component(handle); + } + + public ComponentExport? GetExport(string name) + { + var ret = Native.wasmtime_component_get_export_index(NativeHandle, null, name, (nuint)name.Length); + if (ret == IntPtr.Zero) + return null; + + return new ComponentExport(ret); + } + + public ComponentExport? GetExport(string name, ComponentExport instance_export_index) + { + var ret = Native.wasmtime_component_get_export_index(NativeHandle, instance_export_index.NativeHandle, name, (nuint)name.Length); + if (ret == IntPtr.Zero) + return null; + + return new ComponentExport(ret); + } + + internal class Handle + : SafeHandleZeroOrMinusOneIsInvalid + { + public Handle(IntPtr handle) + : base(true) + { + SetHandle(handle); + } + + protected override bool ReleaseHandle() + { + Native.wasmtime_component_delete(handle); + return true; + } + } + + internal static class Native + { + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_new(Engine.Handle engine, byte* bytes, nuint size, out IntPtr handle); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_delete(IntPtr handle); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_serialize(Handle component, out ByteArray ret); + + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_deserialize(Engine.Handle engine, byte* bytes, nuint size, out IntPtr handle); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_deserialize_file(Engine.Handle engine, string path, out IntPtr handle); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_get_export_index(Handle component, ComponentExport.Handle? instance_export_index, string name, nuint name_len); + } +} \ No newline at end of file diff --git a/src/Components/ComponentExport.cs b/src/Components/ComponentExport.cs new file mode 100644 index 00000000..4fb03259 --- /dev/null +++ b/src/Components/ComponentExport.cs @@ -0,0 +1,57 @@ +using Microsoft.Win32.SafeHandles; +using System; +using System.Runtime.InteropServices; + +namespace Wasmtime.Components; + +public class ComponentExport + : IDisposable +{ + private readonly Handle handle; + + internal Handle NativeHandle + { + get + { + if (handle.IsInvalid || handle.IsClosed) + { + throw new ObjectDisposedException(typeof(Module).FullName); + } + + return handle; + } + } + + internal ComponentExport(IntPtr handle) + { + this.handle = new Handle(handle); + } + + /// + public void Dispose() + { + handle.Dispose(); + } + + internal class Handle + : SafeHandleZeroOrMinusOneIsInvalid + { + public Handle(IntPtr handle) + : base(true) + { + SetHandle(handle); + } + + protected override bool ReleaseHandle() + { + Native.wasmtime_component_export_index_delete(handle); + return true; + } + } + + internal static class Native + { + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_export_index_delete(IntPtr /* wasmtime_component_export_index_t* */ export_index); + } +} \ No newline at end of file diff --git a/src/Components/ComponentFunction.cs b/src/Components/ComponentFunction.cs new file mode 100644 index 00000000..7e390e0d --- /dev/null +++ b/src/Components/ComponentFunction.cs @@ -0,0 +1,19 @@ +namespace Wasmtime.Components; + +/// +/// Represents a Wasmtime function. +/// +public class ComponentFunction +{ + //todo: everything! + + + internal static class Native + { + // [DllImport(Engine.LibraryName)] + //todo: wasmtime_error_t * wasmtime_component_func_call (const wasmtime_component_func_t *func, wasmtime_context_t *context, const wasmtime_component_val_t *args, size_t args_size, wasmtime_component_val_t *results, size_t results_size) + + // [DllImport(Engine.LibraryName)] + //todo: wasmtime_error_t * wasmtime_component_func_post_return (const wasmtime_component_func_t *func, wasmtime_context_t *context) + } +} \ No newline at end of file diff --git a/src/Components/ComponentInstance.cs b/src/Components/ComponentInstance.cs new file mode 100644 index 00000000..00f19632 --- /dev/null +++ b/src/Components/ComponentInstance.cs @@ -0,0 +1,16 @@ +namespace Wasmtime.Components; + +public class ComponentInstance +{ + //todo: everything! + + + internal static class Native + { + //[DllImport(Engine.LibraryName)] + //public static extern IntPtr /* wasmtime_component_export_index_t* */ wasmtime_component_instance_get_export_index (wasmtime_component_instance_t *instance, wasmtime_context_t *context, ComponentExport.Handle instance_export_index, string name, nuint name_len) + + // [DllImport(Engine.LibraryName)] + //todo: bool wasmtime_component_instance_get_func (const wasmtime_component_instance_t *instance, wasmtime_context_t *context, const wasmtime_component_export_index_t *export_index, wasmtime_component_func_t *func_out) + } +} \ No newline at end of file diff --git a/src/Components/ComponentLinker.cs b/src/Components/ComponentLinker.cs new file mode 100644 index 00000000..7f3047e3 --- /dev/null +++ b/src/Components/ComponentLinker.cs @@ -0,0 +1,75 @@ +using Microsoft.Win32.SafeHandles; +using System; +using System.Runtime.InteropServices; + +namespace Wasmtime.Components; + +public class ComponentLinker + : IDisposable +{ + private readonly Handle handle; + + internal Handle NativeHandle + { + get + { + if (handle.IsInvalid || handle.IsClosed) + { + throw new ObjectDisposedException(typeof(Module).FullName); + } + + return handle; + } + } + + internal ComponentLinker(IntPtr handle) + { + this.handle = new Handle(handle); + } + + /// + public void Dispose() + { + handle.Dispose(); + } + + internal class Handle + : SafeHandleZeroOrMinusOneIsInvalid + { + public Handle(IntPtr handle) + : base(true) + { + SetHandle(handle); + } + + protected override bool ReleaseHandle() + { + Native.wasmtime_component_linker_delete(handle); + return true; + } + } + + internal static class Native + { + // [DllImport(Engine.LibraryName)] + //todo: wasmtime_component_linker_t * wasmtime_component_linker_new (const wasm_engine_t *engine) + + // [DllImport(Engine.LibraryName)] + //todo: wasmtime_component_linker_instance_t * wasmtime_component_linker_root (wasmtime_component_linker_t *linker) + + // [DllImport(Engine.LibraryName)] + //todo: wasmtime_error_t * wasmtime_component_linker_instantiate (const wasmtime_component_linker_t *linker, wasmtime_context_t *context, const wasmtime_component_t *component, wasmtime_component_instance_t *instance_out) + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_linker_delete(IntPtr /* wasmtime_component_linker_t* */ linker); + + //todo: wasmtime_error_t * wasmtime_component_linker_instance_add_instance (wasmtime_component_linker_instance_t *linker_instance, const char *name, size_t name_len, wasmtime_component_linker_instance_t **linker_instance_out) + //todo: wasmtime_error_t* wasmtime_component_linker_instance_add_module(wasmtime_component_linker_instance_t* linker_instance, const char* name, size_t name_len, const wasmtime_module_t* module) + //todo: wasmtime_error_t * wasmtime_component_linker_instance_add_func (wasmtime_component_linker_instance_t *linker_instance, const char *name, size_t name_len, wasmtime_component_func_callback_t callback, void *data, void(*finalizer)(void *)) + //todo: wasmtime_error_t * wasmtime_component_linker_add_wasip2 (wasmtime_component_linker_t *linker) + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_linker_instance_delete(IntPtr /* wasmtime_component_linker_instance_t* */ linker_instance); + + } +} \ No newline at end of file diff --git a/src/Components/ComponentValue.cs b/src/Components/ComponentValue.cs new file mode 100644 index 00000000..c4b8f440 --- /dev/null +++ b/src/Components/ComponentValue.cs @@ -0,0 +1,130 @@ +using System.Drawing; +using System.Runtime.InteropServices; +using System.Security.Cryptography; +using static System.Runtime.InteropServices.JavaScript.JSType; + +namespace Wasmtime.Components; + +// todo: everything here: https://docs.wasmtime.dev/c-api/component_2val_8h.html +/* + wasmtime_component_vallist + A vec of a struct wasmtime_component_val + + wasmtime_component_valrecord + A vec of a struct wasmtime_component_valrecord_entry + + wasmtime_component_valtuple + A vec of a struct wasmtime_component_val + + wasmtime_component_valflags + A vec of a wasm_name_t + + wasmtime_component_valvariant_t + Represents a variant type + + wasmtime_component_valresult_t + Represents a result type + + wasmtime_component_valunion_t + Represents possible runtime values which a component function can either consume or produce + + wasmtime_component_val + Represents possible runtime values which a component function can either consume or produce + + wasmtime_component_valrecord_entry + A pair of a name and a value that represents one entry in a value with kind WASMTIME_COMPONENT_RECORD +*/ + +internal enum ComponentValueKind +{ + Bool = 0, + S8 = 1, + U8 = 2, + S16 = 3, + U16 = 4, + S32 = 5, + U32 = 6, + S64 = 7, + U64 = 8, + F32 = 9, + F64 = 10, + Char = 11, + String = 12, + List = 13, + Record = 14, + Tuple = 15, + Variant = 16, + Enum = 17, + Option = 18, + Result = 19, + Flags = 20, +} + +internal static class ComponentValueNative +{ + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_vallist_new(wasmtime_component_vallist_t*out, size_t size, struct wasmtime_component_val * ptr) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_vallist_new_empty (wasmtime_component_vallist_t*out) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_vallist_new_uninit (wasmtime_component_vallist_t*out, size_t size) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_vallist_copy (wasmtime_component_vallist_t* dst, const wasmtime_component_vallist_t* src) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_vallist_delete (wasmtime_component_vallist_t* value) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valrecord_new (wasmtime_component_valrecord_t*out, size_t size, struct wasmtime_component_valrecord_entry * ptr) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valrecord_new_empty (wasmtime_component_valrecord_t*out) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valrecord_new_uninit (wasmtime_component_valrecord_t*out, size_t size) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valrecord_copy (wasmtime_component_valrecord_t* dst, const wasmtime_component_valrecord_t* src) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valrecord_delete (wasmtime_component_valrecord_t* value) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valtuple_new (wasmtime_component_valtuple_t*out, size_t size, struct wasmtime_component_val * ptr) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valtuple_new_empty (wasmtime_component_valtuple_t*out) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valtuple_new_uninit (wasmtime_component_valtuple_t*out, size_t size) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valtuple_copy (wasmtime_component_valtuple_t* dst, const wasmtime_component_valtuple_t* src) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valtuple_delete (wasmtime_component_valtuple_t* value) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valflags_new (wasmtime_component_valflags_t*out, size_t size, wasm_name_t* ptr) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valflags_new_empty (wasmtime_component_valflags_t*out) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valflags_new_uninit (wasmtime_component_valflags_t*out, size_t size) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valflags_copy (wasmtime_component_valflags_t* dst, const wasmtime_component_valflags_t* src) + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_valflags_delete (wasmtime_component_valflags_t* value) + + //[DllImport(Engine.LibraryName)] + //public static extern wasmtime_component_val_t * wasmtime_component_val_new () + + //[DllImport(Engine.LibraryName)] + //public static extern void wasmtime_component_val_delete(wasmtime_component_val_t* value) +} \ No newline at end of file diff --git a/src/Module.cs b/src/Module.cs index 421558e7..6df18def 100644 --- a/src/Module.cs +++ b/src/Module.cs @@ -1,12 +1,15 @@ +using Microsoft.Win32.SafeHandles; using System; using System.Collections.Generic; using System.IO; using System.Runtime.InteropServices; using System.Text; -using Microsoft.Win32.SafeHandles; namespace Wasmtime { + /// + /// Equivalent to wasm_byte_vec_t + /// [StructLayout(LayoutKind.Sequential)] internal unsafe struct ByteArray : IDisposable { @@ -18,6 +21,20 @@ public void Dispose() Native.wasm_byte_vec_delete(this); } + public Span AsSpan() + { + return new Span(data, checked((int)size)); + } + + public byte[] ToArray() + { + var src = AsSpan(); + var dst = new byte[src.Length]; + src.CopyTo(dst); + + return dst; + } + private static class Native { [DllImport(Engine.LibraryName)] @@ -236,22 +253,14 @@ public static Module FromTextStream(Engine engine, string name, Stream stream) /// Returns the serialized module as an array of bytes. public byte[] Serialize() { - var error = Native.wasmtime_module_serialize(this.handle, out var array); + var error = Native.wasmtime_module_serialize(handle, out var bytes); if (error != IntPtr.Zero) { throw WasmtimeException.FromOwnedError(error); } - using (array) - { - var len = checked((int)array.size); - var bytes = new byte[len]; - unsafe - { - Marshal.Copy((IntPtr)array.data, bytes, 0, len); - } - return bytes; - } + using (bytes) + return bytes.ToArray(); } /// diff --git a/src/Wasmtime.csproj b/src/Wasmtime.csproj index b88a7d93..34702165 100644 --- a/src/Wasmtime.csproj +++ b/src/Wasmtime.csproj @@ -157,26 +157,11 @@ The .NET embedding of Wasmtime enables .NET code to instantiate WebAssembly modu - + - - + + From 49ea67e983308dc0eabb8c4fd765d603b30afb9d Mon Sep 17 00:00:00 2001 From: Martin Evans Date: Sun, 17 Aug 2025 21:43:31 +0100 Subject: [PATCH 2/9] Ignoring user settings --- .gitignore | 1 + 1 file changed, 1 insertion(+) diff --git a/.gitignore b/.gitignore index 09fc0eab..c3f91418 100644 --- a/.gitignore +++ b/.gitignore @@ -9,3 +9,4 @@ bin/ obj/ BenchmarkDotNet.Artifacts/ +/Wasmtime.sln.DotSettings.user From 6c2fb31b2360ac0f7dbcbdd7de4d79876263aac6 Mon Sep 17 00:00:00 2001 From: Martin Evans Date: Wed, 20 Aug 2025 16:34:25 +0100 Subject: [PATCH 3/9] Remoed unnecessary imports --- src/Components/ComponentValue.cs | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/src/Components/ComponentValue.cs b/src/Components/ComponentValue.cs index c4b8f440..12d3ece2 100644 --- a/src/Components/ComponentValue.cs +++ b/src/Components/ComponentValue.cs @@ -1,9 +1,4 @@ -using System.Drawing; -using System.Runtime.InteropServices; -using System.Security.Cryptography; -using static System.Runtime.InteropServices.JavaScript.JSType; - -namespace Wasmtime.Components; +namespace Wasmtime.Components; // todo: everything here: https://docs.wasmtime.dev/c-api/component_2val_8h.html /* From 0b3382dc1800f7ad8afe97115df4aae85591f30d Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Thu, 17 Sep 2026 14:40:19 +1200 Subject: [PATCH 4/9] Implement ComponentValueMarshaller for native encoding and decoding of component values - Added ComponentValueMarshaller class to handle conversion between ComponentValue and native wasmtime_component_val_t. - Implemented allocation management for native memory in AllocationScope. - Created methods for writing and reading various component value types, including scalars, strings, lists, tuples, records, options, and results. - Introduced tests for validating the layout and round-trip encoding of component values against Wasmtime's definitions. - Added comprehensive unit tests to ensure correct behavior of the marshaller and its handling of different component value types. --- src/Components/ComponentValue.cs | 555 +++++++++++++++++---- src/Components/ComponentValueMarshaller.cs | 372 ++++++++++++++ tests/ComponentValueLayoutTests.cs | 300 +++++++++++ tests/ComponentValueTests.cs | 443 ++++++++++++++++ 4 files changed, 1572 insertions(+), 98 deletions(-) create mode 100644 src/Components/ComponentValueMarshaller.cs create mode 100644 tests/ComponentValueLayoutTests.cs create mode 100644 tests/ComponentValueTests.cs diff --git a/src/Components/ComponentValue.cs b/src/Components/ComponentValue.cs index 12d3ece2..3b4eee7f 100644 --- a/src/Components/ComponentValue.cs +++ b/src/Components/ComponentValue.cs @@ -1,125 +1,484 @@ -namespace Wasmtime.Components; - -// todo: everything here: https://docs.wasmtime.dev/c-api/component_2val_8h.html -/* - wasmtime_component_vallist - A vec of a struct wasmtime_component_val - - wasmtime_component_valrecord - A vec of a struct wasmtime_component_valrecord_entry - - wasmtime_component_valtuple - A vec of a struct wasmtime_component_val - - wasmtime_component_valflags - A vec of a wasm_name_t - - wasmtime_component_valvariant_t - Represents a variant type - - wasmtime_component_valresult_t - Represents a result type - - wasmtime_component_valunion_t - Represents possible runtime values which a component function can either consume or produce - - wasmtime_component_val - Represents possible runtime values which a component function can either consume or produce - - wasmtime_component_valrecord_entry - A pair of a name and a value that represents one entry in a value with kind WASMTIME_COMPONENT_RECORD -*/ - -internal enum ComponentValueKind +using System; +using System.Collections.Generic; +using System.Globalization; +using System.Runtime.InteropServices; + +namespace Wasmtime.Components; + +/// +/// The type discriminant of a , mirroring +/// wasmtime_component_valkind_t. +/// +public enum ComponentValueKind : byte { + /// A bool. Bool = 0, + /// A signed 8-bit integer. S8 = 1, + /// An unsigned 8-bit integer. U8 = 2, + /// A signed 16-bit integer. S16 = 3, + /// An unsigned 16-bit integer. U16 = 4, + /// A signed 32-bit integer. S32 = 5, + /// An unsigned 32-bit integer. U32 = 6, + /// A signed 64-bit integer. S64 = 7, + /// An unsigned 64-bit integer. U64 = 8, + /// A 32-bit IEEE-754 float. F32 = 9, + /// A 64-bit IEEE-754 float. F64 = 10, + /// A Unicode scalar value. Not yet supported by this binding. Char = 11, + /// A UTF-8 string. String = 12, + /// A homogeneous list. List = 13, + /// A record of named fields. Record = 14, + /// A tuple of positional values. Tuple = 15, + /// A variant. Not yet supported by this binding. Variant = 16, + /// An enumeration, identified by case name. Enum = 17, + /// An optional value. Option = 18, + /// A result value. Result = 19, + /// A set of flags. Not yet supported by this binding. Flags = 20, + /// A resource handle. Not yet supported by this binding. + Resource = 21, + /// A map. Not yet supported by this binding. + Map = 22, } -internal static class ComponentValueNative +/// +/// A value which a component function can consume or produce, mirroring +/// wasmtime_component_val_t. +/// +/// +/// Create values with the static factory methods and read them back with the typed accessors, +/// each of which throws if the value has a different . The +/// , , +/// , and +/// kinds are not yet supported. +/// +public sealed class ComponentValue { - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_vallist_new(wasmtime_component_vallist_t*out, size_t size, struct wasmtime_component_val * ptr) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_vallist_new_empty (wasmtime_component_vallist_t*out) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_vallist_new_uninit (wasmtime_component_vallist_t*out, size_t size) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_vallist_copy (wasmtime_component_vallist_t* dst, const wasmtime_component_vallist_t* src) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_vallist_delete (wasmtime_component_vallist_t* value) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valrecord_new (wasmtime_component_valrecord_t*out, size_t size, struct wasmtime_component_valrecord_entry * ptr) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valrecord_new_empty (wasmtime_component_valrecord_t*out) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valrecord_new_uninit (wasmtime_component_valrecord_t*out, size_t size) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valrecord_copy (wasmtime_component_valrecord_t* dst, const wasmtime_component_valrecord_t* src) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valrecord_delete (wasmtime_component_valrecord_t* value) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valtuple_new (wasmtime_component_valtuple_t*out, size_t size, struct wasmtime_component_val * ptr) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valtuple_new_empty (wasmtime_component_valtuple_t*out) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valtuple_new_uninit (wasmtime_component_valtuple_t*out, size_t size) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valtuple_copy (wasmtime_component_valtuple_t* dst, const wasmtime_component_valtuple_t* src) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valtuple_delete (wasmtime_component_valtuple_t* value) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valflags_new (wasmtime_component_valflags_t*out, size_t size, wasm_name_t* ptr) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valflags_new_empty (wasmtime_component_valflags_t*out) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valflags_new_uninit (wasmtime_component_valflags_t*out, size_t size) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valflags_copy (wasmtime_component_valflags_t* dst, const wasmtime_component_valflags_t* src) - - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_valflags_delete (wasmtime_component_valflags_t* value) - - //[DllImport(Engine.LibraryName)] - //public static extern wasmtime_component_val_t * wasmtime_component_val_new () + private ComponentValue(ComponentValueKind kind) + { + Kind = kind; + } + + /// + /// Gets the type discriminant of this value. + /// + public ComponentValueKind Kind { get; } + + /// Backing store for the integral kinds and for . + internal long Integer { get; private set; } + + /// Backing store for the floating-point kinds. + internal double Real { get; private set; } + + /// String contents, or the case name of an enum. + internal string? Text { get; private set; } + + /// Elements of a list or tuple. + internal IReadOnlyList Items { get; private set; } = Array.Empty(); + + /// Fields of a record, in declaration order. + internal IReadOnlyList> Fields { get; private set; } = + Array.Empty>(); + + /// True for some options and ok results. + internal bool Flag { get; private set; } + + /// + /// Gets the payload of an or + /// , or null when there is none. + /// + public ComponentValue? Payload { get; private set; } + + /// + /// Gets whether this carries a value. + /// + public bool IsSome => Expect(ComponentValueKind.Option, Flag); + + /// + /// Gets whether this is the ok case. + /// + public bool IsOk => Expect(ComponentValueKind.Result, Flag); + + /// Creates a bool value. + /// The value. + /// The component value. + public static ComponentValue Bool(bool value) => + new ComponentValue(ComponentValueKind.Bool) { Integer = value ? 1 : 0 }; + + /// Creates a signed 8-bit integer value. + /// The value. + /// The component value. + public static ComponentValue S8(sbyte value) => + new ComponentValue(ComponentValueKind.S8) { Integer = value }; + + /// Creates an unsigned 8-bit integer value. + /// The value. + /// The component value. + public static ComponentValue U8(byte value) => + new ComponentValue(ComponentValueKind.U8) { Integer = value }; + + /// Creates a signed 16-bit integer value. + /// The value. + /// The component value. + public static ComponentValue S16(short value) => + new ComponentValue(ComponentValueKind.S16) { Integer = value }; + + /// Creates an unsigned 16-bit integer value. + /// The value. + /// The component value. + public static ComponentValue U16(ushort value) => + new ComponentValue(ComponentValueKind.U16) { Integer = value }; + + /// Creates a signed 32-bit integer value. + /// The value. + /// The component value. + public static ComponentValue S32(int value) => + new ComponentValue(ComponentValueKind.S32) { Integer = value }; + + /// Creates an unsigned 32-bit integer value. + /// The value. + /// The component value. + public static ComponentValue U32(uint value) => + new ComponentValue(ComponentValueKind.U32) { Integer = value }; + + /// Creates a signed 64-bit integer value. + /// The value. + /// The component value. + public static ComponentValue S64(long value) => + new ComponentValue(ComponentValueKind.S64) { Integer = value }; + + /// Creates an unsigned 64-bit integer value. + /// The value. + /// The component value. + public static ComponentValue U64(ulong value) => + new ComponentValue(ComponentValueKind.U64) { Integer = unchecked((long)value) }; + + /// Creates a 32-bit float value. + /// The value. + /// The component value. + public static ComponentValue F32(float value) => + new ComponentValue(ComponentValueKind.F32) { Real = value }; + + /// Creates a 64-bit float value. + /// The value. + /// The component value. + public static ComponentValue F64(double value) => + new ComponentValue(ComponentValueKind.F64) { Real = value }; + + /// Creates a string value. + /// The value. + /// The component value. + /// Thrown if is null. + public static ComponentValue String(string value) + { + if (value is null) + { + throw new ArgumentNullException(nameof(value)); + } + + return new ComponentValue(ComponentValueKind.String) { Text = value }; + } + + /// Creates an enumeration value. + /// The name of the enum case, as declared in WIT. + /// The component value. + /// Thrown if is null. + /// Component model enums travel by case name rather than by ordinal. + public static ComponentValue Enum(string caseName) + { + if (caseName is null) + { + throw new ArgumentNullException(nameof(caseName)); + } + + return new ComponentValue(ComponentValueKind.Enum) { Text = caseName }; + } + + /// Creates a list value. + /// The elements of the list, which are copied. + /// The component value. + /// Thrown if is null. + public static ComponentValue List(IReadOnlyList items) + { + if (items is null) + { + throw new ArgumentNullException(nameof(items)); + } + + return new ComponentValue(ComponentValueKind.List) { Items = Copy(items) }; + } + + /// Creates a tuple value. + /// The elements of the tuple, which are copied. + /// The component value. + /// Thrown if is null. + public static ComponentValue Tuple(IReadOnlyList items) + { + if (items is null) + { + throw new ArgumentNullException(nameof(items)); + } + + return new ComponentValue(ComponentValueKind.Tuple) { Items = Copy(items) }; + } + + /// Creates a record value. + /// The fields of the record in declaration order, which are copied. + /// The component value. + /// Thrown if is null. + public static ComponentValue Record(IReadOnlyList> fields) + { + if (fields is null) + { + throw new ArgumentNullException(nameof(fields)); + } + + return new ComponentValue(ComponentValueKind.Record) { Fields = Copy(fields) }; + } + + // The Owned* factories take ownership instead of copying, for lists nothing else references. + internal static ComponentValue OwnedList(List items) => + new ComponentValue(ComponentValueKind.List) { Items = items }; + + internal static ComponentValue OwnedTuple(List items) => + new ComponentValue(ComponentValueKind.Tuple) { Items = items }; + + internal static ComponentValue OwnedRecord(List> fields) => + new ComponentValue(ComponentValueKind.Record) { Fields = fields }; + + private static T[] Copy(IReadOnlyList source) + { + if (source.Count == 0) + { + return Array.Empty(); + } + + var copy = new T[source.Count]; + for (var i = 0; i < copy.Length; i++) + { + copy[i] = source[i]; + } + + return copy; + } + + /// Creates an option value carrying a payload. + /// The payload. + /// The component value. + /// Thrown if is null. + public static ComponentValue Some(ComponentValue value) + { + if (value is null) + { + throw new ArgumentNullException(nameof(value)); + } + + return new ComponentValue(ComponentValueKind.Option) { Payload = value, Flag = true }; + } + + /// Creates an empty option value. + /// The component value. + public static ComponentValue None() => + new ComponentValue(ComponentValueKind.Option) { Flag = false }; + + /// Creates the ok case of a result. + /// The payload, or null for a result with no ok type. + /// The component value. + public static ComponentValue Ok(ComponentValue? value = null) => + new ComponentValue(ComponentValueKind.Result) { Payload = value, Flag = true }; + + /// Creates the err case of a result. + /// The payload, or null for a result with no err type. + /// The component value. + public static ComponentValue Err(ComponentValue? value = null) => + new ComponentValue(ComponentValueKind.Result) { Payload = value, Flag = false }; + + /// Gets the value as a bool. + /// The value. + public bool AsBool() => Expect(ComponentValueKind.Bool, Integer != 0); + + /// Gets the value as a signed 8-bit integer. + /// The value. + public sbyte AsS8() => Expect(ComponentValueKind.S8, unchecked((sbyte)Integer)); + + /// Gets the value as an unsigned 8-bit integer. + /// The value. + public byte AsU8() => Expect(ComponentValueKind.U8, unchecked((byte)Integer)); + + /// Gets the value as a signed 16-bit integer. + /// The value. + public short AsS16() => Expect(ComponentValueKind.S16, unchecked((short)Integer)); + + /// Gets the value as an unsigned 16-bit integer. + /// The value. + public ushort AsU16() => Expect(ComponentValueKind.U16, unchecked((ushort)Integer)); + + /// Gets the value as a signed 32-bit integer. + /// The value. + public int AsS32() => Expect(ComponentValueKind.S32, unchecked((int)Integer)); + + /// Gets the value as an unsigned 32-bit integer. + /// The value. + public uint AsU32() => Expect(ComponentValueKind.U32, unchecked((uint)Integer)); + + /// Gets the value as a signed 64-bit integer. + /// The value. + public long AsS64() => Expect(ComponentValueKind.S64, Integer); + + /// Gets the value as an unsigned 64-bit integer. + /// The value. + public ulong AsU64() => Expect(ComponentValueKind.U64, unchecked((ulong)Integer)); + + /// Gets the value as a 32-bit float. + /// The value. + public float AsF32() => Expect(ComponentValueKind.F32, (float)Real); + + /// Gets the value as a 64-bit float. + /// The value. + public double AsF64() => Expect(ComponentValueKind.F64, Real); + + /// Gets the value as a string. + /// The value. + public string AsString() => Expect(ComponentValueKind.String, Text ?? string.Empty); + + /// Gets the case name of an enumeration value. + /// The case name. + public string AsEnum() => Expect(ComponentValueKind.Enum, Text ?? string.Empty); + + /// Gets the elements of a list. + /// The elements. + public IReadOnlyList AsList() => Expect(ComponentValueKind.List, Items); + + /// Gets the elements of a tuple. + /// The elements. + public IReadOnlyList AsTuple() => Expect(ComponentValueKind.Tuple, Items); + + /// Gets the fields of a record. + /// The fields, in declaration order. + public IReadOnlyList> AsRecord() => + Expect(ComponentValueKind.Record, Fields); + + /// + /// Gets the value of a record field by name. + /// + /// The field name. + /// The field's value. + /// + /// Thrown if this is not a record, or if it has no such field. + /// + public ComponentValue Field(string name) + { + if (!TryGetField(name, out var value)) + { + throw new InvalidOperationException($"Record has no field '{name}'."); + } + + return value!; + } + + /// + /// Gets the value of a record field by name, if present. + /// + /// The field name. + /// The field's value, or null if there is no such field. + /// True if the field was found. + /// Thrown if this is not a record. + public bool TryGetField(string name, out ComponentValue? value) + { + foreach (var field in AsRecord()) + { + if (field.Key == name) + { + value = field.Value; + return true; + } + } + + value = null; + return false; + } + + private T Expect(ComponentValueKind expected, T value) + { + if (Kind != expected) + { + throw new InvalidOperationException($"Expected component value kind {expected}, got {Kind}."); + } + + return value; + } + + /// + public override string ToString() + { + switch (Kind) + { + case ComponentValueKind.Bool: + return (Integer != 0).ToString(CultureInfo.InvariantCulture); + case ComponentValueKind.F32: + case ComponentValueKind.F64: + return Real.ToString(CultureInfo.InvariantCulture); + case ComponentValueKind.U64: + return unchecked((ulong)Integer).ToString(CultureInfo.InvariantCulture); + case ComponentValueKind.String: + return $"\"{Text}\""; + case ComponentValueKind.Enum: + return Text ?? string.Empty; + case ComponentValueKind.List: + case ComponentValueKind.Tuple: + return $"[{string.Join(", ", Items)}]"; + case ComponentValueKind.Record: + var fields = new string[Fields.Count]; + for (var i = 0; i < Fields.Count; i++) + { + fields[i] = $"{Fields[i].Key}: {Fields[i].Value}"; + } + + return $"{{{string.Join(", ", fields)}}}"; + case ComponentValueKind.Option: + return Flag ? $"some({Payload})" : "none"; + case ComponentValueKind.Result: + return Flag ? $"ok({Payload})" : $"err({Payload})"; + default: + return Integer.ToString(CultureInfo.InvariantCulture); + } + } +} - //[DllImport(Engine.LibraryName)] - //public static extern void wasmtime_component_val_delete(wasmtime_component_val_t* value) +internal static class ComponentValueNative +{ + /// + /// Performs a deep copy of into . The contents + /// of are owned by Wasmtime and must be released with + /// . + /// + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_val_clone( + IntPtr /* const wasmtime_component_val_t* */ src, + IntPtr /* wasmtime_component_val_t* */ dst); + + /// + /// Deallocates the memory owned by , but not the storage of + /// itself. Only valid for embedder-owned storage. + /// + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_val_delete(IntPtr /* wasmtime_component_val_t* */ value); } \ No newline at end of file diff --git a/src/Components/ComponentValueMarshaller.cs b/src/Components/ComponentValueMarshaller.cs new file mode 100644 index 00000000..705a8e40 --- /dev/null +++ b/src/Components/ComponentValueMarshaller.cs @@ -0,0 +1,372 @@ +using System; +using System.Collections.Generic; +using System.Runtime.InteropServices; +using System.Text; + +namespace Wasmtime.Components; + +/// +/// Converts between and the native wasmtime_component_val_t +/// layout. +/// +/// +/// +/// Ownership is asymmetric. Arguments are declared const by the C API, so Wasmtime never +/// takes them: everything written here is allocated by us and released by the caller's +/// . Results are the opposite, being allocated by Wasmtime's own +/// allocator, so their contents must be released with +/// wasmtime_component_val_delete rather than freed directly. Mixing the two allocators +/// corrupts the heap. +/// +/// +/// The native structures are written field by field rather than through the +/// wasmtime_component_vallist_new family of helpers, because a value must be built in +/// place inside an array element or a record entry. +/// +/// +internal static class ComponentValueMarshaller +{ + /// Size of wasmtime_component_val_t: a 1-byte kind padded to 8, plus a 24-byte union. + public const int ValueSize = 32; + + /// Offset of the of union within wasmtime_component_val_t. + public const int ValuePayloadOffset = 8; + + /// Size of wasmtime_component_valrecord_entry_t: a 16-byte name plus a 32-byte value. + public const int RecordEntrySize = 48; + + /// Offset of the value within wasmtime_component_valrecord_entry_t. + public const int RecordEntryValueOffset = 16; + + /// Offset of the data pointer within a vec or a wasm_name_t. + private const int VectorDataOffset = 8; + + /// + /// Tracks the native allocations made while writing arguments, so they can be released + /// together once the call has returned. + /// + public sealed class AllocationScope : IDisposable + { + private readonly List allocations = new List(); + private bool disposed; + + ~AllocationScope() + { + FreeAllocations(); + } + + /// + /// Allocates zeroed native memory whose lifetime is bound to this scope. + /// + /// The number of bytes to allocate. + /// A pointer to the allocation. + public IntPtr Allocate(int bytes) + { + if (disposed) + { + throw new ObjectDisposedException(nameof(AllocationScope)); + } + + var pointer = Marshal.AllocHGlobal(bytes); + allocations.Add(pointer); + + unsafe + { + new Span((void*)pointer, bytes).Clear(); + } + + return pointer; + } + + /// + public void Dispose() + { + if (disposed) + { + return; + } + + disposed = true; + FreeAllocations(); + GC.SuppressFinalize(this); + } + + private void FreeAllocations() + { + foreach (var pointer in allocations) + { + Marshal.FreeHGlobal(pointer); + } + + allocations.Clear(); + } + } + + /// + /// Writes a value into caller-owned native storage. + /// + /// The value to write. + /// A pointer to bytes of storage. + /// The scope owning any secondary allocations the value needs. + public static void Write(ComponentValue value, IntPtr destination, AllocationScope scope) + { + Marshal.WriteByte(destination, (byte)value.Kind); + var payload = destination + ValuePayloadOffset; + + switch (value.Kind) + { + case ComponentValueKind.Bool: + Marshal.WriteByte(payload, (byte)(value.Integer != 0 ? 1 : 0)); + break; + + case ComponentValueKind.S8: + case ComponentValueKind.U8: + Marshal.WriteByte(payload, unchecked((byte)value.Integer)); + break; + + case ComponentValueKind.S16: + case ComponentValueKind.U16: + Marshal.WriteInt16(payload, unchecked((short)value.Integer)); + break; + + case ComponentValueKind.S32: + case ComponentValueKind.U32: + Marshal.WriteInt32(payload, unchecked((int)value.Integer)); + break; + + case ComponentValueKind.F32: + Marshal.WriteInt32(payload, Extensions.SingleToInt32Bits((float)value.Real)); + break; + + case ComponentValueKind.S64: + case ComponentValueKind.U64: + Marshal.WriteInt64(payload, value.Integer); + break; + + case ComponentValueKind.F64: + Marshal.WriteInt64(payload, BitConverter.DoubleToInt64Bits(value.Real)); + break; + + case ComponentValueKind.String: + case ComponentValueKind.Enum: + WriteName(value.Text ?? string.Empty, payload, scope); + break; + + case ComponentValueKind.List: + case ComponentValueKind.Tuple: + WriteVector(value.Items, payload, scope); + break; + + case ComponentValueKind.Record: + WriteRecord(value.Fields, payload, scope); + break; + + case ComponentValueKind.Option: + Marshal.WriteIntPtr(payload, value.Flag ? WriteBoxed(value.Payload!, scope) : IntPtr.Zero); + break; + + case ComponentValueKind.Result: + Marshal.WriteByte(payload, (byte)(value.Flag ? 1 : 0)); + Marshal.WriteIntPtr( + payload + VectorDataOffset, + value.Payload is null ? IntPtr.Zero : WriteBoxed(value.Payload, scope)); + break; + + default: + throw new NotSupportedException( + $"Writing component values of kind {value.Kind} is not supported."); + } + } + + /// + /// Reads a value from native storage. + /// + /// A pointer to a wasmtime_component_val_t. + /// The value that was read. + public static ComponentValue Read(IntPtr source) + { + var kind = (ComponentValueKind)Marshal.ReadByte(source); + var payload = source + ValuePayloadOffset; + + switch (kind) + { + case ComponentValueKind.Bool: + return ComponentValue.Bool(Marshal.ReadByte(payload) != 0); + + case ComponentValueKind.S8: + return ComponentValue.S8(unchecked((sbyte)Marshal.ReadByte(payload))); + + case ComponentValueKind.U8: + return ComponentValue.U8(Marshal.ReadByte(payload)); + + case ComponentValueKind.S16: + return ComponentValue.S16(Marshal.ReadInt16(payload)); + + case ComponentValueKind.U16: + return ComponentValue.U16(unchecked((ushort)Marshal.ReadInt16(payload))); + + case ComponentValueKind.S32: + return ComponentValue.S32(Marshal.ReadInt32(payload)); + + case ComponentValueKind.U32: + return ComponentValue.U32(unchecked((uint)Marshal.ReadInt32(payload))); + + case ComponentValueKind.S64: + return ComponentValue.S64(Marshal.ReadInt64(payload)); + + case ComponentValueKind.U64: + return ComponentValue.U64(unchecked((ulong)Marshal.ReadInt64(payload))); + + case ComponentValueKind.F32: + return ComponentValue.F32(Extensions.Int32BitsToSingle(Marshal.ReadInt32(payload))); + + case ComponentValueKind.F64: + return ComponentValue.F64(BitConverter.Int64BitsToDouble(Marshal.ReadInt64(payload))); + + case ComponentValueKind.String: + return ComponentValue.String(ReadName(payload)); + + case ComponentValueKind.Enum: + return ComponentValue.Enum(ReadName(payload)); + + case ComponentValueKind.List: + return ComponentValue.OwnedList(ReadVector(payload)); + + case ComponentValueKind.Tuple: + return ComponentValue.OwnedTuple(ReadVector(payload)); + + case ComponentValueKind.Record: + return ComponentValue.OwnedRecord(ReadRecord(payload)); + + case ComponentValueKind.Option: + var some = Marshal.ReadIntPtr(payload); + return some == IntPtr.Zero ? ComponentValue.None() : ComponentValue.Some(Read(some)); + + case ComponentValueKind.Result: + var isOk = Marshal.ReadByte(payload) != 0; + var inner = Marshal.ReadIntPtr(payload + VectorDataOffset); + var result = inner == IntPtr.Zero ? null : Read(inner); + return isOk ? ComponentValue.Ok(result) : ComponentValue.Err(result); + + default: + throw new NotSupportedException( + $"Reading component values of kind {kind} is not supported."); + } + } + + private static IntPtr WriteBoxed(ComponentValue value, AllocationScope scope) + { + var boxed = scope.Allocate(ValueSize); + Write(value, boxed, scope); + return boxed; + } + + private static void WriteName(string text, IntPtr destination, AllocationScope scope) + { + var bytes = Encoding.UTF8.GetBytes(text); + + // A zero-length allocation would still need a non-null pointer, so always take at least one byte. + var buffer = scope.Allocate(Math.Max(bytes.Length, 1)); + Marshal.Copy(bytes, 0, buffer, bytes.Length); + + Marshal.WriteIntPtr(destination, (IntPtr)bytes.Length); + Marshal.WriteIntPtr(destination + VectorDataOffset, buffer); + } + + private static void WriteVector(IReadOnlyList items, IntPtr destination, AllocationScope scope) + { + // Read once: the buffer size, loop bound and written length must agree even if the list changes. + var count = items.Count; + + // checked: an overflowing size would otherwise under-allocate and be written past. + var buffer = scope.Allocate(Math.Max(checked(count * ValueSize), 1)); + var element = buffer; + + for (var i = 0; i < count; i++) + { + Write(items[i], element, scope); + element += ValueSize; + } + + Marshal.WriteIntPtr(destination, (IntPtr)count); + Marshal.WriteIntPtr(destination + VectorDataOffset, buffer); + } + + private static void WriteRecord( + IReadOnlyList> fields, + IntPtr destination, + AllocationScope scope) + { + var count = fields.Count; + var buffer = scope.Allocate(Math.Max(checked(count * RecordEntrySize), 1)); + var entry = buffer; + + for (var i = 0; i < count; i++) + { + WriteName(fields[i].Key, entry, scope); + Write(fields[i].Value, entry + RecordEntryValueOffset, scope); + entry += RecordEntrySize; + } + + Marshal.WriteIntPtr(destination, (IntPtr)count); + Marshal.WriteIntPtr(destination + VectorDataOffset, buffer); + } + + private static string ReadName(IntPtr source) + { + var size = ReadLength(source); + var data = Marshal.ReadIntPtr(source + VectorDataOffset); + + return data == IntPtr.Zero || size == 0 ? string.Empty : Extensions.PtrToStringUTF8(data, size); + } + + private static List ReadVector(IntPtr source) + { + var size = ReadLength(source); + var element = Marshal.ReadIntPtr(source + VectorDataOffset); + var items = new List(size); + + for (var i = 0; i < size; i++) + { + items.Add(Read(element)); + element += ValueSize; + } + + return items; + } + + private static List> ReadRecord(IntPtr source) + { + var size = ReadLength(source); + var entry = Marshal.ReadIntPtr(source + VectorDataOffset); + var fields = new List>(size); + + for (var i = 0; i < size; i++) + { + fields.Add(new KeyValuePair( + ReadName(entry), + Read(entry + RecordEntryValueOffset))); + + entry += RecordEntrySize; + } + + return fields; + } + + /// + /// Reads a native size_t length, rejecting values too large to be represented + /// managed-side rather than silently truncating them. + /// + private static int ReadLength(IntPtr source) + { + // Reinterpreted rather than converted, because size_t is unsigned. + var size = unchecked((ulong)Marshal.ReadIntPtr(source).ToInt64()); + if (size > int.MaxValue) + { + throw new NotSupportedException( + $"A component value with {size} elements exceeds the maximum supported length of {int.MaxValue}."); + } + + return (int)size; + } +} diff --git a/tests/ComponentValueLayoutTests.cs b/tests/ComponentValueLayoutTests.cs new file mode 100644 index 00000000..c77ae501 --- /dev/null +++ b/tests/ComponentValueLayoutTests.cs @@ -0,0 +1,300 @@ +using System; +using System.Collections.Generic; +using System.Runtime.InteropServices; +using FluentAssertions; +using Wasmtime.Components; +using Xunit; + +namespace Wasmtime.Tests +{ + /// + /// Validates the native encoding of component values against Wasmtime itself. + /// + /// + /// + /// only proves that our writer and reader agree with each + /// other; it would still pass if every offset were wrong in the same way. These tests close + /// that gap by routing each value through wasmtime_component_val_clone, which walks + /// the structure using Wasmtime's own definition of the layout: + /// + /// + /// we encode a value into source; + /// Wasmtime deep-copies it, reading source with its layout; + /// we decode destination, which Wasmtime wrote with its layout. + /// + /// + /// A mismatch in any size or offset therefore corrupts the value or crashes, rather than + /// cancelling out. This needs no engine, store or component. + /// + /// + public sealed class ComponentValueLayoutTests + { + private static ComponentValue Clone(ComponentValue value) + { + using var scope = new ComponentValueMarshaller.AllocationScope(); + var source = scope.Allocate(ComponentValueMarshaller.ValueSize); + ComponentValueMarshaller.Write(value, source, scope); + + var destination = Marshal.AllocHGlobal(ComponentValueMarshaller.ValueSize); + try + { + unsafe + { + new Span((void*)destination, ComponentValueMarshaller.ValueSize).Clear(); + } + + ComponentValueNative.wasmtime_component_val_clone(source, destination); + return ComponentValueMarshaller.Read(destination); + } + finally + { + ComponentValueNative.wasmtime_component_val_delete(destination); + Marshal.FreeHGlobal(destination); + } + } + + /// + /// Clones the value through Wasmtime and asserts the result is structurally identical. + /// + private static void Survives(ComponentValue value) + { + AssertEquivalent(value, Clone(value), "$"); + } + + private static void AssertEquivalent(ComponentValue expected, ComponentValue actual, string path) + { + actual.Kind.Should().Be(expected.Kind, "kind at {0}", path); + + switch (expected.Kind) + { + case ComponentValueKind.Bool: + actual.AsBool().Should().Be(expected.AsBool(), "value at {0}", path); + break; + + case ComponentValueKind.S8: + case ComponentValueKind.U8: + case ComponentValueKind.S16: + case ComponentValueKind.U16: + case ComponentValueKind.S32: + case ComponentValueKind.U32: + case ComponentValueKind.S64: + case ComponentValueKind.U64: + actual.Integer.Should().Be(expected.Integer, "value at {0}", path); + break; + + case ComponentValueKind.F32: + case ComponentValueKind.F64: + actual.Real.Should().Be(expected.Real, "value at {0}", path); + break; + + case ComponentValueKind.String: + case ComponentValueKind.Enum: + actual.Text.Should().Be(expected.Text, "text at {0}", path); + break; + + case ComponentValueKind.List: + case ComponentValueKind.Tuple: + actual.Items.Should().HaveCount(expected.Items.Count, "length at {0}", path); + for (var i = 0; i < expected.Items.Count; i++) + { + AssertEquivalent(expected.Items[i], actual.Items[i], $"{path}[{i}]"); + } + + break; + + case ComponentValueKind.Record: + actual.Fields.Should().HaveCount(expected.Fields.Count, "field count at {0}", path); + for (var i = 0; i < expected.Fields.Count; i++) + { + actual.Fields[i].Key.Should().Be(expected.Fields[i].Key, "field name at {0}[{1}]", path, i); + AssertEquivalent( + expected.Fields[i].Value, + actual.Fields[i].Value, + $"{path}.{expected.Fields[i].Key}"); + } + + break; + + case ComponentValueKind.Option: + case ComponentValueKind.Result: + actual.Flag.Should().Be(expected.Flag, "discriminant at {0}", path); + if (expected.Payload is null) + { + actual.Payload.Should().BeNull("payload at {0}", path); + } + else + { + actual.Payload.Should().NotBeNull("payload at {0}", path); + AssertEquivalent(expected.Payload, actual.Payload!, $"{path}.payload"); + } + + break; + + default: + throw new NotSupportedException($"No comparison for kind {expected.Kind}."); + } + } + + [Fact] + public void ScalarsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Bool(true)); + Survives(ComponentValue.Bool(false)); + Survives(ComponentValue.S8(sbyte.MinValue)); + Survives(ComponentValue.S8(-1)); + Survives(ComponentValue.U8(byte.MaxValue)); + Survives(ComponentValue.S16(short.MinValue)); + Survives(ComponentValue.U16(ushort.MaxValue)); + Survives(ComponentValue.S32(int.MinValue)); + Survives(ComponentValue.S32(-1)); + Survives(ComponentValue.U32(uint.MaxValue)); + Survives(ComponentValue.S64(long.MinValue)); + Survives(ComponentValue.U64(ulong.MaxValue)); + Survives(ComponentValue.F32(float.MinValue)); + Survives(ComponentValue.F32(float.NaN)); + Survives(ComponentValue.F64(-2.718281828459045)); + Survives(ComponentValue.F64(double.PositiveInfinity)); + } + + [Fact] + public void StringsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.String(string.Empty)); + Survives(ComponentValue.String("hello")); + Survives(ComponentValue.String("h\u00e9llo \u4e16\u754c \ud83c\udf89")); + Survives(ComponentValue.String(new string('x', 4096))); + } + + [Fact] + public void EnumsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Enum("warning")); + Survives(ComponentValue.Enum("a")); + } + + [Fact] + public void ListsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.List([])); + Survives(ComponentValue.List(new[] { ComponentValue.S32(1) })); + Survives(ComponentValue.List(new[] + { + ComponentValue.S32(1), + ComponentValue.S32(-2), + ComponentValue.S32(int.MaxValue), + })); + Survives(ComponentValue.List(new[] + { + ComponentValue.String("first"), + ComponentValue.String(string.Empty), + ComponentValue.String("third"), + })); + } + + [Fact] + public void NestedListsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.List(new[] + { + ComponentValue.List(new[] { ComponentValue.String("a"), ComponentValue.String("b") }), + ComponentValue.List([]), + ComponentValue.List(new[] { ComponentValue.String("c") }), + })); + } + + [Fact] + public void TuplesSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Tuple(new[] + { + ComponentValue.Bool(true), + ComponentValue.String("two"), + ComponentValue.F64(3.5), + })); + } + + [Fact] + public void RecordsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Record([])); + Survives(ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(36.6)), + new KeyValuePair("timestamp", ComponentValue.U64(1234567890123)), + })); + } + + [Fact] + public void OptionsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.None()); + Survives(ComponentValue.Some(ComponentValue.S32(42))); + Survives(ComponentValue.Some(ComponentValue.String("boxed"))); + Survives(ComponentValue.Some(ComponentValue.Some(ComponentValue.String("x")))); + } + + [Fact] + public void ResultsSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Ok()); + Survives(ComponentValue.Err()); + Survives(ComponentValue.Ok(ComponentValue.S32(7))); + Survives(ComponentValue.Err(ComponentValue.String("boom"))); + } + + [Fact] + public void DeeplyNestedValuesSurviveACloneThroughWasmtime() + { + Survives(ComponentValue.Record(new[] + { + new KeyValuePair("readings", ComponentValue.List(new[] + { + ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(1.5)), + new KeyValuePair("note", ComponentValue.Some(ComponentValue.String("ok"))), + }), + ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(2.5)), + new KeyValuePair("note", ComponentValue.None()), + }), + })), + new KeyValuePair("level", ComponentValue.Enum("warning")), + new KeyValuePair("outcome", ComponentValue.Ok(ComponentValue.Tuple(new[] + { + ComponentValue.U16(200), + ComponentValue.String("done"), + }))), + })); + } + + /// + /// A large list exercises the element stride, which a single-element list cannot: an + /// over-estimated still round-trips one + /// element correctly but misaligns every element after the first. + /// + [Fact] + public void ElementStrideSurvivesACloneThroughWasmtime() + { + var items = new ComponentValue[64]; + for (var i = 0; i < items.Length; i++) + { + items[i] = ComponentValue.Record(new[] + { + new KeyValuePair("index", ComponentValue.S32(i)), + new KeyValuePair("name", ComponentValue.String($"item-{i}")), + }); + } + + var result = Clone(ComponentValue.List(items)).AsList(); + + result.Should().HaveCount(items.Length); + for (var i = 0; i < items.Length; i++) + { + result[i].Field("index").AsS32().Should().Be(i); + result[i].Field("name").AsString().Should().Be($"item-{i}"); + } + } + } +} diff --git a/tests/ComponentValueTests.cs b/tests/ComponentValueTests.cs new file mode 100644 index 00000000..70f0cb46 --- /dev/null +++ b/tests/ComponentValueTests.cs @@ -0,0 +1,443 @@ +using System; +using System.Collections.Generic; +using FluentAssertions; +using Wasmtime.Components; +using Xunit; + +namespace Wasmtime.Tests +{ + /// + /// Round-trip tests for the component value marshaller. + /// + /// + /// These verify that the writer and the reader agree on the native encoding. They do not + /// prove that the encoding matches Wasmtime's, which is fixed by the struct sizes and + /// offsets taken from wasmtime/component/val.h and is exercised end to end once a + /// component function is actually called. + /// + public sealed class ComponentValueTests + { + private static ComponentValue RoundTrip(ComponentValue value) + { + using var scope = new ComponentValueMarshaller.AllocationScope(); + var pointer = scope.Allocate(ComponentValueMarshaller.ValueSize); + ComponentValueMarshaller.Write(value, pointer, scope); + return ComponentValueMarshaller.Read(pointer); + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public void ItRoundTripsBool(bool value) + { + var result = RoundTrip(ComponentValue.Bool(value)); + + result.Kind.Should().Be(ComponentValueKind.Bool); + result.AsBool().Should().Be(value); + } + + [Theory] + [InlineData(sbyte.MinValue)] + [InlineData((sbyte)-1)] + [InlineData((sbyte)0)] + [InlineData(sbyte.MaxValue)] + public void ItRoundTripsS8(sbyte value) + { + RoundTrip(ComponentValue.S8(value)).AsS8().Should().Be(value); + } + + [Theory] + [InlineData(byte.MinValue)] + [InlineData((byte)1)] + [InlineData(byte.MaxValue)] + public void ItRoundTripsU8(byte value) + { + RoundTrip(ComponentValue.U8(value)).AsU8().Should().Be(value); + } + + [Theory] + [InlineData(short.MinValue)] + [InlineData((short)-1)] + [InlineData(short.MaxValue)] + public void ItRoundTripsS16(short value) + { + RoundTrip(ComponentValue.S16(value)).AsS16().Should().Be(value); + } + + [Theory] + [InlineData(ushort.MinValue)] + [InlineData(ushort.MaxValue)] + public void ItRoundTripsU16(ushort value) + { + RoundTrip(ComponentValue.U16(value)).AsU16().Should().Be(value); + } + + [Theory] + [InlineData(int.MinValue)] + [InlineData(-1)] + [InlineData(0)] + [InlineData(int.MaxValue)] + public void ItRoundTripsS32(int value) + { + RoundTrip(ComponentValue.S32(value)).AsS32().Should().Be(value); + } + + [Theory] + [InlineData(uint.MinValue)] + [InlineData(1u)] + [InlineData(uint.MaxValue)] + public void ItRoundTripsU32(uint value) + { + RoundTrip(ComponentValue.U32(value)).AsU32().Should().Be(value); + } + + [Theory] + [InlineData(long.MinValue)] + [InlineData(-1L)] + [InlineData(long.MaxValue)] + public void ItRoundTripsS64(long value) + { + RoundTrip(ComponentValue.S64(value)).AsS64().Should().Be(value); + } + + [Theory] + [InlineData(ulong.MinValue)] + [InlineData(1ul)] + [InlineData(ulong.MaxValue)] + public void ItRoundTripsU64(ulong value) + { + RoundTrip(ComponentValue.U64(value)).AsU64().Should().Be(value); + } + + [Theory] + [InlineData(0.0f)] + [InlineData(3.14159f)] + [InlineData(float.MinValue)] + [InlineData(float.MaxValue)] + [InlineData(float.Epsilon)] + [InlineData(float.NaN)] + [InlineData(float.PositiveInfinity)] + [InlineData(float.NegativeInfinity)] + public void ItRoundTripsF32(float value) + { + RoundTrip(ComponentValue.F32(value)).AsF32().Should().Be(value); + } + + /// + /// -0.0 == 0.0, so the sign of zero is only observable through the bits. + /// + [Fact] + public void ItPreservesTheSignOfZero() + { + var f32 = RoundTrip(ComponentValue.F32(-0.0f)).AsF32(); + BitConverter.SingleToInt32Bits(f32).Should().Be(BitConverter.SingleToInt32Bits(-0.0f)); + + var f64 = RoundTrip(ComponentValue.F64(-0.0)).AsF64(); + BitConverter.DoubleToInt64Bits(f64).Should().Be(BitConverter.DoubleToInt64Bits(-0.0)); + } + + [Theory] + [InlineData(0.0)] + [InlineData(-2.718281828459045)] + [InlineData(double.MinValue)] + [InlineData(double.MaxValue)] + [InlineData(double.Epsilon)] + [InlineData(double.NaN)] + [InlineData(double.PositiveInfinity)] + [InlineData(double.NegativeInfinity)] + public void ItRoundTripsF64(double value) + { + RoundTrip(ComponentValue.F64(value)).AsF64().Should().Be(value); + } + + [Theory] + [InlineData("")] + [InlineData("hello")] + [InlineData("h\u00e9llo \u4e16\u754c \ud83c\udf89")] + [InlineData("embedded\0null")] + public void ItRoundTripsString(string value) + { + var result = RoundTrip(ComponentValue.String(value)); + + result.Kind.Should().Be(ComponentValueKind.String); + result.AsString().Should().Be(value); + } + + [Fact] + public void ItRoundTripsEnum() + { + var result = RoundTrip(ComponentValue.Enum("warning")); + + result.Kind.Should().Be(ComponentValueKind.Enum); + result.AsEnum().Should().Be("warning"); + } + + [Fact] + public void ItRoundTripsEmptyList() + { + RoundTrip(ComponentValue.List([])).AsList().Should().BeEmpty(); + } + + [Fact] + public void ItRoundTripsList() + { + var value = ComponentValue.List(new[] + { + ComponentValue.S32(1), + ComponentValue.S32(-2), + ComponentValue.S32(int.MaxValue), + }); + + var result = RoundTrip(value).AsList(); + + result.Should().HaveCount(3); + result[0].AsS32().Should().Be(1); + result[1].AsS32().Should().Be(-2); + result[2].AsS32().Should().Be(int.MaxValue); + } + + [Fact] + public void ItRoundTripsNestedLists() + { + var value = ComponentValue.List(new[] + { + ComponentValue.List(new[] { ComponentValue.String("a"), ComponentValue.String("b") }), + ComponentValue.List([]), + ComponentValue.List(new[] { ComponentValue.String("c") }), + }); + + var result = RoundTrip(value).AsList(); + + result.Should().HaveCount(3); + result[0].AsList().Should().HaveCount(2); + result[0].AsList()[1].AsString().Should().Be("b"); + result[1].AsList().Should().BeEmpty(); + result[2].AsList()[0].AsString().Should().Be("c"); + } + + [Fact] + public void ItRoundTripsTupleOfMixedKinds() + { + var value = ComponentValue.Tuple(new[] + { + ComponentValue.Bool(true), + ComponentValue.String("two"), + ComponentValue.F64(3.5), + }); + + var result = RoundTrip(value); + + result.Kind.Should().Be(ComponentValueKind.Tuple); + result.AsTuple()[0].AsBool().Should().BeTrue(); + result.AsTuple()[1].AsString().Should().Be("two"); + result.AsTuple()[2].AsF64().Should().Be(3.5); + } + + [Fact] + public void ItRoundTripsRecordPreservingFieldOrder() + { + var value = ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(36.6)), + new KeyValuePair("timestamp", ComponentValue.U64(1234567890123)), + }); + + var result = RoundTrip(value); + + result.Kind.Should().Be(ComponentValueKind.Record); + result.AsRecord().Should().HaveCount(2); + result.AsRecord()[0].Key.Should().Be("value"); + result.AsRecord()[1].Key.Should().Be("timestamp"); + result.Field("value").AsF64().Should().Be(36.6); + result.Field("timestamp").AsU64().Should().Be(1234567890123); + } + + [Fact] + public void ItRoundTripsEmptyRecord() + { + RoundTrip(ComponentValue.Record([])) + .AsRecord().Should().BeEmpty(); + } + + [Fact] + public void ItRoundTripsSomeAndNone() + { + var some = RoundTrip(ComponentValue.Some(ComponentValue.S32(42))); + some.Kind.Should().Be(ComponentValueKind.Option); + some.IsSome.Should().BeTrue(); + some.Payload!.AsS32().Should().Be(42); + + var none = RoundTrip(ComponentValue.None()); + none.Kind.Should().Be(ComponentValueKind.Option); + none.IsSome.Should().BeFalse(); + none.Payload.Should().BeNull(); + } + + [Fact] + public void ItRoundTripsNestedOption() + { + var result = RoundTrip(ComponentValue.Some(ComponentValue.Some(ComponentValue.String("x")))); + + result.IsSome.Should().BeTrue(); + result.Payload!.IsSome.Should().BeTrue(); + result.Payload!.Payload!.AsString().Should().Be("x"); + } + + [Fact] + public void ItRoundTripsResultWithPayload() + { + var ok = RoundTrip(ComponentValue.Ok(ComponentValue.S32(7))); + ok.Kind.Should().Be(ComponentValueKind.Result); + ok.IsOk.Should().BeTrue(); + ok.Payload!.AsS32().Should().Be(7); + + var err = RoundTrip(ComponentValue.Err(ComponentValue.String("boom"))); + err.IsOk.Should().BeFalse(); + err.Payload!.AsString().Should().Be("boom"); + } + + [Fact] + public void ItRoundTripsResultWithoutPayload() + { + var ok = RoundTrip(ComponentValue.Ok()); + ok.IsOk.Should().BeTrue(); + ok.Payload.Should().BeNull(); + + var err = RoundTrip(ComponentValue.Err()); + err.IsOk.Should().BeFalse(); + err.Payload.Should().BeNull(); + } + + [Fact] + public void ItRoundTripsDeeplyNestedValues() + { + var value = ComponentValue.Record(new[] + { + new KeyValuePair("readings", ComponentValue.List(new[] + { + ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(1.5)), + new KeyValuePair("note", ComponentValue.Some(ComponentValue.String("ok"))), + }), + ComponentValue.Record(new[] + { + new KeyValuePair("value", ComponentValue.F64(2.5)), + new KeyValuePair("note", ComponentValue.None()), + }), + })), + new KeyValuePair("level", ComponentValue.Enum("warning")), + }); + + var result = RoundTrip(value); + var readings = result.Field("readings").AsList(); + + readings.Should().HaveCount(2); + readings[0].Field("value").AsF64().Should().Be(1.5); + readings[0].Field("note").Payload!.AsString().Should().Be("ok"); + readings[1].Field("value").AsF64().Should().Be(2.5); + readings[1].Field("note").IsSome.Should().BeFalse(); + result.Field("level").AsEnum().Should().Be("warning"); + } + + [Fact] + public void ItThrowsReadingAnUnsupportedKind() + { + using var scope = new ComponentValueMarshaller.AllocationScope(); + var pointer = scope.Allocate(ComponentValueMarshaller.ValueSize); + System.Runtime.InteropServices.Marshal.WriteByte(pointer, (byte)ComponentValueKind.Char); + + var act = () => ComponentValueMarshaller.Read(pointer); + + act.Should().Throw().WithMessage("*Char*"); + } + + /// + /// A native length is a size_t, so it can exceed what a managed length can hold. + /// Truncating it would produce a negative or far-too-small length and read out of bounds, + /// so it has to be rejected instead. + /// + [Theory] + [InlineData(ComponentValueKind.String)] + [InlineData(ComponentValueKind.List)] + [InlineData(ComponentValueKind.Record)] + public void ItRejectsALengthTooLargeToRepresent(ComponentValueKind kind) + { + using var scope = new ComponentValueMarshaller.AllocationScope(); + var pointer = scope.Allocate(ComponentValueMarshaller.ValueSize); + + System.Runtime.InteropServices.Marshal.WriteByte(pointer, (byte)kind); + + // All bits set reads back as a huge unsigned size on both 32- and 64-bit. + System.Runtime.InteropServices.Marshal.WriteIntPtr( + pointer + ComponentValueMarshaller.ValuePayloadOffset, new IntPtr(-1)); + + // Non-null so the guard, rather than a null check, is what rejects this. + System.Runtime.InteropServices.Marshal.WriteIntPtr( + pointer + ComponentValueMarshaller.ValuePayloadOffset + 8, pointer); + + var act = () => ComponentValueMarshaller.Read(pointer); + + act.Should().Throw() + .WithMessage("*exceeds the maximum supported length*"); + } + + [Fact] + public void ItThrowsAccessingTheWrongKind() + { + var value = ComponentValue.S32(1); + + value.Invoking(v => v.AsString()).Should().Throw(); + value.Invoking(v => v.AsRecord()).Should().Throw(); + value.Invoking(v => v.IsOk).Should().Throw(); + } + + [Fact] + public void ItCopiesTheCallersCollections() + { + var items = new List { ComponentValue.S32(1) }; + var fields = new List> { new("a", ComponentValue.S32(1)) }; + + var list = ComponentValue.List(items); + var tuple = ComponentValue.Tuple(items); + var record = ComponentValue.Record(fields); + + items.Add(ComponentValue.S32(2)); + fields.Add(new("b", ComponentValue.S32(2))); + + list.AsList().Should().ContainSingle(); + tuple.AsTuple().Should().ContainSingle(); + record.AsRecord().Should().ContainSingle(); + } + + [Fact] + public void ItThrowsForAMissingRecordField() + { + var value = ComponentValue.Record(new[] + { + new KeyValuePair("present", ComponentValue.S32(1)), + }); + + value.TryGetField("present", out var found).Should().BeTrue(); + found!.AsS32().Should().Be(1); + + value.TryGetField("absent", out var missing).Should().BeFalse(); + missing.Should().BeNull(); + + value.Invoking(v => v.Field("absent")) + .Should().Throw() + .WithMessage("*absent*"); + } + + [Fact] + public void ItRejectsNullArguments() + { + FluentActions.Invoking(() => ComponentValue.String(null!)).Should().Throw(); + FluentActions.Invoking(() => ComponentValue.Enum(null!)).Should().Throw(); + FluentActions.Invoking(() => ComponentValue.List(null!)).Should().Throw(); + FluentActions.Invoking(() => ComponentValue.Tuple(null!)).Should().Throw(); + FluentActions.Invoking(() => ComponentValue.Record(null!)).Should().Throw(); + FluentActions.Invoking(() => ComponentValue.Some(null!)).Should().Throw(); + } + } +} From 98b0a03830bfa26ba6e3ff1c81388b9ffbb059a0 Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Thu, 17 Sep 2026 21:20:14 +1200 Subject: [PATCH 5/9] Add component linker host-function support --- src/Components/Component.cs | 46 ++- src/Components/ComponentInstance.cs | 31 +- src/Components/ComponentLinker.cs | 164 +++++++++- src/Components/ComponentLinkerInstance.cs | 304 ++++++++++++++++++ src/Components/ComponentValueMarshaller.cs | 21 ++ tests/ComponentFixture.cs | 43 +++ tests/ComponentLinkerTests.cs | 191 +++++++++++ tests/ComponentValueLayoutTests.cs | 48 +++ tests/Components/host-import-declarations.wat | 6 + tests/Components/host-import.wat | 30 ++ tests/Components/tiny.wat | 12 + tests/Wasmtime.Tests.csproj | 1 + 12 files changed, 874 insertions(+), 23 deletions(-) create mode 100644 src/Components/ComponentLinkerInstance.cs create mode 100644 tests/ComponentFixture.cs create mode 100644 tests/ComponentLinkerTests.cs create mode 100644 tests/Components/host-import-declarations.wat create mode 100644 tests/Components/host-import.wat create mode 100644 tests/Components/tiny.wat diff --git a/src/Components/Component.cs b/src/Components/Component.cs index 1dfcd573..840f1b26 100644 --- a/src/Components/Component.cs +++ b/src/Components/Component.cs @@ -1,11 +1,12 @@ -using Microsoft.Win32.SafeHandles; -using System; +using System; using System.Runtime.InteropServices; +using System.Text; +using Microsoft.Win32.SafeHandles; namespace Wasmtime.Components; /// -/// Representation of a component in the component model. +/// Representation of a component in the component model. /// public class Component : IDisposable @@ -65,7 +66,44 @@ public static Component FromBytes(Engine engine, ReadOnlySpan bytes) } /// - /// This function serializes compiled component artifacts as blob data. + /// Creates a based on a WebAssembly text format representation. + /// + /// The engine to use for the component. + /// The WebAssembly text format representation of the component. + /// Returns a new . + public static Component FromText(Engine engine, string text) + { + if (engine is null) + { + throw new ArgumentNullException(nameof(engine)); + } + + if (text is null) + { + throw new ArgumentNullException(nameof(text)); + } + + unsafe + { + var textBytes = Encoding.UTF8.GetBytes(text); + fixed (byte* ptr = textBytes) + { + var error = Module.Native.wasmtime_wat2wasm(ptr, (nuint)textBytes.Length, out var componentBytes); + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + using (componentBytes) + { + return FromBytes(engine, new ReadOnlySpan(componentBytes.data, checked((int)componentBytes.size))); + } + } + } + } + + /// + /// This function serializes compiled component artifacts as blob data. /// /// If the conversion is successful, the serialized compiled component. public byte[] Serialize() diff --git a/src/Components/ComponentInstance.cs b/src/Components/ComponentInstance.cs index 00f19632..5e991448 100644 --- a/src/Components/ComponentInstance.cs +++ b/src/Components/ComponentInstance.cs @@ -1,12 +1,39 @@ -namespace Wasmtime.Components; +using System.Runtime.InteropServices; +namespace Wasmtime.Components; + +/// +/// An instantiated . +/// +/// +/// Instances are identified by an integer within their store and have no destructor, so this +/// type is not disposable. Passing an instance to a different than the one +/// it was created in may abort the process. +/// public class ComponentInstance { - //todo: everything! + internal readonly Native.Instance instance; + internal readonly Store store; + internal ComponentInstance(Store store, Native.Instance instance) + { + this.store = store; + this.instance = instance; + } internal static class Native { + /// Mirrors wasmtime_component_instance_t. + [StructLayout(LayoutKind.Explicit, Size = 16)] + internal struct Instance + { + [FieldOffset(0)] + public ulong StoreId; + + [FieldOffset(8)] + public uint Private; + } + //[DllImport(Engine.LibraryName)] //public static extern IntPtr /* wasmtime_component_export_index_t* */ wasmtime_component_instance_get_export_index (wasmtime_component_instance_t *instance, wasmtime_context_t *context, ComponentExport.Handle instance_export_index, string name, nuint name_len) diff --git a/src/Components/ComponentLinker.cs b/src/Components/ComponentLinker.cs index 7f3047e3..604fde6f 100644 --- a/src/Components/ComponentLinker.cs +++ b/src/Components/ComponentLinker.cs @@ -1,13 +1,38 @@ -using Microsoft.Win32.SafeHandles; -using System; +using System; using System.Runtime.InteropServices; +using Microsoft.Win32.SafeHandles; namespace Wasmtime.Components; +/// +/// Used to instantiate a , providing the host functions and other +/// definitions that satisfy its imports. +/// public class ComponentLinker : IDisposable { private readonly Handle handle; + private ComponentLinkerInstance? root; + + /// + /// Creates a new for the given engine. + /// + /// The engine to create the linker for. + /// Thrown if is null. + public ComponentLinker(Engine engine) + { + if (engine is null) + { + throw new ArgumentNullException(nameof(engine)); + } + + handle = new Handle(Native.wasmtime_component_linker_new(engine.NativeHandle)); + } + + internal ComponentLinker(IntPtr handle) + { + this.handle = new Handle(handle); + } internal Handle NativeHandle { @@ -15,21 +40,120 @@ internal Handle NativeHandle { if (handle.IsInvalid || handle.IsClosed) { - throw new ObjectDisposedException(typeof(Module).FullName); + throw new ObjectDisposedException(typeof(ComponentLinker).FullName); + } + + if (root is not null) + { + throw new InvalidOperationException( + "The linker cannot be used while the linker instance returned by Root() is still in use."); } return handle; } } - internal ComponentLinker(IntPtr handle) + /// + /// Sets whether later definitions are allowed to shadow previous ones. + /// + /// True to allow shadowing. + public bool AllowShadowing { - this.handle = new Handle(handle); + set + { + Native.wasmtime_component_linker_allow_shadowing(NativeHandle, value); + } + } + + /// + /// Returns the root instance of this linker, used to define names into the root namespace. + /// + /// The root instance. The linker cannot be used again until this is disposed. + public ComponentLinkerInstance Root() + { + var current = NativeHandle; + root = new ComponentLinkerInstance( + Native.wasmtime_component_linker_root(current), + () => root = null); + + return root; + } + + /// + /// Adds all WASI Preview 2 interfaces to this linker. + /// + /// + /// The store used for instantiation must have a WASI configuration set with + /// . + /// + public void AddWasiPreview2() + { + var error = Native.wasmtime_component_linker_add_wasip2(NativeHandle); + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + } + + /// + /// Defines every import of the given component that is not already defined as a function + /// which traps when called. + /// + /// The component whose imports should be satisfied. + /// Thrown if is null. + public void DefineUnknownImportsAsTraps(Component component) + { + if (component is null) + { + throw new ArgumentNullException(nameof(component)); + } + + var error = Native.wasmtime_component_linker_define_unknown_imports_as_traps( + NativeHandle, component.NativeHandle); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + } + + /// + /// Instantiates a component in the given store. + /// + /// The store to instantiate the component in. + /// The component to instantiate. + /// The instantiated component. + /// Thrown if an argument is null. + /// Thrown if an import could not be satisfied. + public ComponentInstance Instantiate(Store store, Component component) + { + if (store is null) + { + throw new ArgumentNullException(nameof(store)); + } + + if (component is null) + { + throw new ArgumentNullException(nameof(component)); + } + + var error = Native.wasmtime_component_linker_instantiate( + NativeHandle, store.Context.handle, component.NativeHandle, out var instance); + + GC.KeepAlive(store); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + return new ComponentInstance(store, instance); } /// public void Dispose() { + root?.Dispose(); handle.Dispose(); } @@ -51,25 +175,31 @@ protected override bool ReleaseHandle() internal static class Native { - // [DllImport(Engine.LibraryName)] - //todo: wasmtime_component_linker_t * wasmtime_component_linker_new (const wasm_engine_t *engine) + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_linker_new(Engine.Handle engine); - // [DllImport(Engine.LibraryName)] - //todo: wasmtime_component_linker_instance_t * wasmtime_component_linker_root (wasmtime_component_linker_t *linker) + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_linker_allow_shadowing( + Handle linker, [MarshalAs(UnmanagedType.I1)] bool allow); - // [DllImport(Engine.LibraryName)] - //todo: wasmtime_error_t * wasmtime_component_linker_instantiate (const wasmtime_component_linker_t *linker, wasmtime_context_t *context, const wasmtime_component_t *component, wasmtime_component_instance_t *instance_out) + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_linker_root(Handle linker); [DllImport(Engine.LibraryName)] - public static extern void wasmtime_component_linker_delete(IntPtr /* wasmtime_component_linker_t* */ linker); + public static extern IntPtr wasmtime_component_linker_instantiate( + Handle linker, + IntPtr context, + Component.Handle component, + out ComponentInstance.Native.Instance instance_out); - //todo: wasmtime_error_t * wasmtime_component_linker_instance_add_instance (wasmtime_component_linker_instance_t *linker_instance, const char *name, size_t name_len, wasmtime_component_linker_instance_t **linker_instance_out) - //todo: wasmtime_error_t* wasmtime_component_linker_instance_add_module(wasmtime_component_linker_instance_t* linker_instance, const char* name, size_t name_len, const wasmtime_module_t* module) - //todo: wasmtime_error_t * wasmtime_component_linker_instance_add_func (wasmtime_component_linker_instance_t *linker_instance, const char *name, size_t name_len, wasmtime_component_func_callback_t callback, void *data, void(*finalizer)(void *)) - //todo: wasmtime_error_t * wasmtime_component_linker_add_wasip2 (wasmtime_component_linker_t *linker) + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_linker_define_unknown_imports_as_traps( + Handle linker, Component.Handle component); [DllImport(Engine.LibraryName)] - public static extern void wasmtime_component_linker_instance_delete(IntPtr /* wasmtime_component_linker_instance_t* */ linker_instance); + public static extern IntPtr wasmtime_component_linker_add_wasip2(Handle linker); + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_linker_delete(IntPtr /* wasmtime_component_linker_t* */ linker); } } \ No newline at end of file diff --git a/src/Components/ComponentLinkerInstance.cs b/src/Components/ComponentLinkerInstance.cs new file mode 100644 index 00000000..85df42ef --- /dev/null +++ b/src/Components/ComponentLinkerInstance.cs @@ -0,0 +1,304 @@ +using System; +using System.Runtime.InteropServices; +using System.Text; +using Microsoft.Win32.SafeHandles; + +namespace Wasmtime.Components; + +/// +/// A callback implementing a component function defined by the host. +/// +/// The arguments passed by the guest. +/// +/// The results to return to the guest. Every element must be assigned. +/// +public delegate void ComponentFunctionCallback( + ReadOnlySpan arguments, + Span results); + +/// +/// An instance being defined within a , used to define names into +/// a namespace. +/// +/// +/// Obtaining one of these acquires exclusive access to whatever it came from: neither the owning +/// linker nor a parent instance may be used until this is disposed. This type enforces that by +/// throwing rather than allowing the undefined behaviour the C API warns about. +/// +public sealed class ComponentLinkerInstance : IDisposable +{ + private readonly Handle handle; + private readonly Action onDisposed; + private ComponentLinkerInstance? child; + private bool disposed; + + internal ComponentLinkerInstance(IntPtr handle, Action onDisposed) + { + this.handle = new Handle(handle); + this.onDisposed = onDisposed; + } + + internal Handle NativeHandle + { + get + { + if (handle.IsInvalid || handle.IsClosed) + { + throw new ObjectDisposedException(typeof(ComponentLinkerInstance).FullName); + } + + if (child is not null) + { + throw new InvalidOperationException( + "This linker instance cannot be used while a nested instance created from it is still in use."); + } + + return handle; + } + } + + /// + /// Defines a nested instance within this instance. + /// + /// The name of the nested instance. + /// The nested instance, which must be disposed before this one is used again. + /// Thrown if is null. + public ComponentLinkerInstance AddInstance(string name) + { + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + var current = NativeHandle; + var nameBytes = Encoding.UTF8.GetBytes(name); + + unsafe + { + fixed (byte* namePtr = nameBytes) + { + var error = Native.wasmtime_component_linker_instance_add_instance( + current, namePtr, (nuint)nameBytes.Length, out var nested); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + child = new ComponentLinkerInstance(nested, () => child = null); + return child; + } + } + } + + /// + /// Defines a core WebAssembly module within this instance. + /// + /// The name to define the module as. + /// The module. + /// Thrown if an argument is null. + public void AddModule(string name, Module module) + { + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + if (module is null) + { + throw new ArgumentNullException(nameof(module)); + } + + var current = NativeHandle; + var nameBytes = Encoding.UTF8.GetBytes(name); + + unsafe + { + fixed (byte* namePtr = nameBytes) + { + var error = Native.wasmtime_component_linker_instance_add_module( + current, namePtr, (nuint)nameBytes.Length, module.NativeHandle); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + } + } + } + + /// + /// Defines a function within this instance. + /// + /// The name of the function. + /// The implementation of the function. + /// Thrown if an argument is null. + public void DefineFunction(string name, ComponentFunctionCallback callback) + { + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + if (callback is null) + { + throw new ArgumentNullException(nameof(callback)); + } + + var current = NativeHandle; + var nameBytes = Encoding.UTF8.GetBytes(name); + + Native.ComponentFuncCallback trampoline = (env, context, type, args, nargs, results, nresults) => + Invoke(callback, name, args, (int)nargs, results, (int)nresults); + + unsafe + { + fixed (byte* namePtr = nameBytes) + { + var error = Native.wasmtime_component_linker_instance_add_func( + current, + namePtr, + (nuint)nameBytes.Length, + trampoline, + GCHandle.ToIntPtr(GCHandle.Alloc(trampoline)), + Function.Finalizer); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + } + } + } + + private static IntPtr Invoke( + ComponentFunctionCallback callback, + string name, + IntPtr args, + int argumentCount, + IntPtr results, + int resultCount) + { + try + { + var arguments = argumentCount == 0 ? Array.Empty() : new ComponentValue[argumentCount]; + for (var i = 0; i < argumentCount; i++) + { + arguments[i] = ComponentValueMarshaller.Read(args + (i * ComponentValueMarshaller.ValueSize)); + } + + var produced = resultCount == 0 ? Array.Empty() : new ComponentValue[resultCount]; + callback(arguments, produced); + + for (var i = 0; i < resultCount; i++) + { + if (produced[i] is null) + { + throw new InvalidOperationException( + $"The callback for component function '{name}' did not assign result {i}."); + } + + ComponentValueMarshaller.WriteOwned( + produced[i], + results + (i * ComponentValueMarshaller.ValueSize)); + } + + return IntPtr.Zero; + } + catch (Exception ex) + { + return HandleCallbackException(ex); + } + } + + /// + /// Converts a managed exception into an owned wasmtime_error_t, since an exception + /// must never propagate across the native-to-managed transition. + /// + private static IntPtr HandleCallbackException(Exception ex) + { + try + { + Function.CallbackErrorCause = ex is WasmtimeException wasmtimeException + ? wasmtimeException.InnerException + : ex; + + return Native.wasmtime_error_new(ex.Message); + } + catch (Exception separateException) + { + // See Function.HandleCallbackException: unwinding through native frames is undefined + // behaviour, so failing fast is the only safe option left. + Environment.FailFast(separateException.Message, separateException); + throw; + } + } + + /// + public void Dispose() + { + if (disposed) + { + return; + } + + disposed = true; + + // A nested instance borrows from this one, so it has to go first. + child?.Dispose(); + handle.Dispose(); + onDisposed(); + } + + internal class Handle + : SafeHandleZeroOrMinusOneIsInvalid + { + public Handle(IntPtr handle) + : base(true) + { + SetHandle(handle); + } + + protected override bool ReleaseHandle() + { + Native.wasmtime_component_linker_instance_delete(handle); + return true; + } + } + + internal static class Native + { + public delegate IntPtr ComponentFuncCallback( + IntPtr env, + IntPtr context, + IntPtr type, + IntPtr args, + nuint nargs, + IntPtr results, + nuint nresults); + + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_linker_instance_add_instance( + Handle linker_instance, byte* name, nuint name_len, out IntPtr linker_instance_out); + + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_linker_instance_add_module( + Handle linker_instance, byte* name, nuint name_len, Module.Handle module); + + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_linker_instance_add_func( + Handle linker_instance, + byte* name, + nuint name_len, + ComponentFuncCallback callback, + IntPtr data, + Function.Native.Finalizer? finalizer); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_linker_instance_delete( + IntPtr /* wasmtime_component_linker_instance_t* */ linker_instance); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_error_new([MarshalAs(Extensions.LPUTF8Str)] string message); + } +} diff --git a/src/Components/ComponentValueMarshaller.cs b/src/Components/ComponentValueMarshaller.cs index 705a8e40..98d2c66d 100644 --- a/src/Components/ComponentValueMarshaller.cs +++ b/src/Components/ComponentValueMarshaller.cs @@ -178,6 +178,27 @@ public static void Write(ComponentValue value, IntPtr destination, AllocationSco } } + /// + /// Writes a value into storage whose contents Wasmtime will free, as required for the + /// results of a host-defined function. + /// + /// The value to write. + /// A pointer to bytes of storage. + /// + /// The value is first built with our own allocator and then deep-copied by Wasmtime, so that + /// every heap payload it ends up owning came from its own allocator. The extra copy is the + /// price of not having to mirror each vec constructor. + /// + public static void WriteOwned(ComponentValue value, IntPtr destination) + { + using (var scope = new AllocationScope()) + { + var scratch = scope.Allocate(ValueSize); + Write(value, scratch, scope); + ComponentValueNative.wasmtime_component_val_clone(scratch, destination); + } + } + /// /// Reads a value from native storage. /// diff --git a/tests/ComponentFixture.cs b/tests/ComponentFixture.cs new file mode 100644 index 00000000..9936b7fc --- /dev/null +++ b/tests/ComponentFixture.cs @@ -0,0 +1,43 @@ +using System; +using System.IO; +using Wasmtime.Components; + +namespace Wasmtime.Tests +{ + public abstract class ComponentFixture : IDisposable + { + public ComponentFixture() + { + Engine = new Engine(GetEngineConfig()); + + Store = new Store(Engine); + } + + public virtual Config GetEngineConfig() + { + return new Config() + .WithComponentModel(true); + } + + public Component LoadComponent(string fileName) => + Component.FromText(Engine, File.ReadAllText(Path.Combine("Components", fileName))); + + public void Dispose() + { + if (!(Store is null)) + { + Store.Dispose(); + Store = null; + } + + if (!(Engine is null)) + { + Engine.Dispose(); + Engine = null; + } + } + + public Engine Engine { get; set; } + public Store Store { get; set; } + } +} \ No newline at end of file diff --git a/tests/ComponentLinkerTests.cs b/tests/ComponentLinkerTests.cs new file mode 100644 index 00000000..433e7f90 --- /dev/null +++ b/tests/ComponentLinkerTests.cs @@ -0,0 +1,191 @@ +using System; +using FluentAssertions; +using Wasmtime.Components; +using Xunit; + +namespace Wasmtime.Tests +{ + public class ComponentLinkerFixture : ComponentFixture { } + + public sealed class ComponentLinkerTests : IClassFixture + { + private readonly ComponentLinkerFixture fixture; + + public ComponentLinkerTests(ComponentLinkerFixture fixture) + { + this.fixture = fixture; + } + + [Fact] + public void ItThrowsForANullEngine() + { + FluentActions.Invoking(() => new ComponentLinker(null!)) + .Should().Throw(); + } + + [Fact] + public void ItInstantiatesAComponentWithoutImports() + { + using var linker = new ComponentLinker(fixture.Engine); + using var component = fixture.LoadComponent("tiny.wat"); + + var instance = linker.Instantiate(fixture.Store, component); + + instance.Should().NotBeNull(); + } + + [Fact] + public void ItThrowsWhenAnImportIsNotDefined() + { + using var linker = new ComponentLinker(fixture.Engine); + using var component = fixture.LoadComponent("host-import-declarations.wat"); + + linker.Invoking(l => l.Instantiate(fixture.Store, component)) + .Should().Throw(); + } + + [Fact] + public void ItInstantiatesWithHostFunctionsDefined() + { + using var linker = new ComponentLinker(fixture.Engine); + using var component = fixture.LoadComponent("host-import.wat"); + + var receivedInput = 0; + DefineHost(linker, value => receivedInput = value); + + linker.Instantiate(fixture.Store, component).Should().NotBeNull(); + receivedInput.Should().Be(21); + } + + [Fact] + public void ItSatisfiesUnknownImportsWithTraps() + { + using var linker = new ComponentLinker(fixture.Engine); + using var component = fixture.LoadComponent("host-import-declarations.wat"); + + linker.DefineUnknownImportsAsTraps(component); + + linker.Instantiate(fixture.Store, component).Should().NotBeNull(); + } + + [Fact] + public void ItThrowsWhenUsingTheLinkerWhileTheRootIsAlive() + { + using var linker = new ComponentLinker(fixture.Engine); + using var component = fixture.LoadComponent("tiny.wat"); + + var root = linker.Root(); + + linker.Invoking(l => l.Instantiate(fixture.Store, component)) + .Should().Throw(); + linker.Invoking(l => l.Root()) + .Should().Throw(); + + root.Dispose(); + + linker.Instantiate(fixture.Store, component).Should().NotBeNull(); + } + + [Fact] + public void ItThrowsWhenUsingAnInstanceWhileItsNestedInstanceIsAlive() + { + using var linker = new ComponentLinker(fixture.Engine); + using var root = linker.Root(); + + var nested = root.AddInstance("host"); + + root.Invoking(r => r.AddInstance("other")) + .Should().Throw(); + + nested.Dispose(); + + root.Invoking(r => r.AddInstance("other")).Should().NotThrow(); + } + + [Fact] + public void ItDisposesNestedInstancesWithTheirParent() + { + var linker = new ComponentLinker(fixture.Engine); + var root = linker.Root(); + var nested = root.AddInstance("host"); + + linker.Dispose(); + + nested.Invoking(n => n.DefineFunction("f", (args, results) => { })) + .Should().Throw(); + } + + [Fact] + public void ItThrowsForDuplicateDefinitionsUnlessShadowingIsAllowed() + { + using var linker = new ComponentLinker(fixture.Engine); + + using (var root = linker.Root()) + { + using var host = root.AddInstance("host"); + host.DefineFunction("transform", (args, results) => results[0] = ComponentValue.S32(0)); + + host.Invoking(h => h.DefineFunction("transform", (args, results) => results[0] = ComponentValue.S32(0))) + .Should().Throw(); + } + + linker.AllowShadowing = true; + + using (var root = linker.Root()) + { + using var host = root.AddInstance("host"); + host.Invoking(h => h.DefineFunction("transform", (args, results) => results[0] = ComponentValue.S32(0))) + .Should().NotThrow(); + } + } + + [Fact] + public void ItThrowsForNullArguments() + { + using var linker = new ComponentLinker(fixture.Engine); + + linker.Invoking(l => l.Instantiate(null!, fixture.LoadComponent("tiny.wat"))).Should().Throw(); + linker.Invoking(l => l.Instantiate(fixture.Store, null!)).Should().Throw(); + linker.Invoking(l => l.DefineUnknownImportsAsTraps(null!)).Should().Throw(); + + using var root = linker.Root(); + root.Invoking(r => r.AddInstance(null!)).Should().Throw(); + root.Invoking(r => r.DefineFunction(null!, (args, results) => { })).Should().Throw(); + root.Invoking(r => r.DefineFunction("f", null!)).Should().Throw(); + root.Invoking(r => r.AddModule(null!, null!)).Should().Throw(); + } + + [Fact] + public void ItThrowsAfterDisposal() + { + var linker = new ComponentLinker(fixture.Engine); + linker.Dispose(); + + linker.Invoking(l => l.Root()).Should().Throw(); + } + + [Fact] + public void ItAddsWasiPreview2() + { + using var linker = new ComponentLinker(fixture.Engine); + + linker.Invoking(l => l.AddWasiPreview2()).Should().NotThrow(); + } + + private static void DefineHost(ComponentLinker linker, Action onTransform) + { + using var root = linker.Root(); + using var host = root.AddInstance("host"); + + host.DefineFunction("transform", (arguments, results) => + { + var input = arguments[0].AsS32(); + onTransform(input); + results[0] = ComponentValue.S32(input * 2); + }); + + host.DefineFunction("greet", (arguments, results) => + results[0] = ComponentValue.String($"hello, {arguments[0].AsString()}")); + } + } +} diff --git a/tests/ComponentValueLayoutTests.cs b/tests/ComponentValueLayoutTests.cs index c77ae501..54a37439 100644 --- a/tests/ComponentValueLayoutTests.cs +++ b/tests/ComponentValueLayoutTests.cs @@ -269,6 +269,54 @@ public void DeeplyNestedValuesSurviveACloneThroughWasmtime() })); } + /// + /// Host function results are written with + /// and then freed by Wasmtime, so every heap payload must come from Wasmtime's allocator. + /// Freeing memory allocated by the wrong allocator corrupts the heap, so this deletes the + /// value after reading it back. + /// + [Fact] + public void OwnedWritesUseWasmtimesAllocator() + { + var values = new[] + { + ComponentValue.S32(7), + ComponentValue.String("a host-allocated string"), + ComponentValue.String(string.Empty), + ComponentValue.Enum("warning"), + ComponentValue.List(new[] { ComponentValue.String("a"), ComponentValue.String("b") }), + ComponentValue.Record(new[] + { + new KeyValuePair("name", ComponentValue.String("x")), + new KeyValuePair("value", ComponentValue.F64(1.5)), + }), + ComponentValue.Some(ComponentValue.String("boxed")), + ComponentValue.None(), + ComponentValue.Ok(ComponentValue.S32(1)), + ComponentValue.Err(ComponentValue.String("boom")), + }; + + foreach (var value in values) + { + var destination = Marshal.AllocHGlobal(ComponentValueMarshaller.ValueSize); + try + { + unsafe + { + new Span((void*)destination, ComponentValueMarshaller.ValueSize).Clear(); + } + + ComponentValueMarshaller.WriteOwned(value, destination); + AssertEquivalent(value, ComponentValueMarshaller.Read(destination), "$"); + } + finally + { + ComponentValueNative.wasmtime_component_val_delete(destination); + Marshal.FreeHGlobal(destination); + } + } + } + /// /// A large list exercises the element stride, which a single-element list cannot: an /// over-estimated still round-trips one diff --git a/tests/Components/host-import-declarations.wat b/tests/Components/host-import-declarations.wat new file mode 100644 index 00000000..0a959e60 --- /dev/null +++ b/tests/Components/host-import-declarations.wat @@ -0,0 +1,6 @@ +;; Component importing host functions without invoking them. +(component + (import "host" (instance $h + (export "transform" (func (param "x" s32) (result s32))) + (export "greet" (func (param "name" string) (result string))))) +) diff --git a/tests/Components/host-import.wat b/tests/Components/host-import.wat new file mode 100644 index 00000000..cd222dcd --- /dev/null +++ b/tests/Components/host-import.wat @@ -0,0 +1,30 @@ +;; Component importing a "host" instance, used to test host function definitions. +;; +;; The core start function calls "transform" during instantiation, exercising both callback +;; argument and result handling. "greet" is only imported, so it must still be defined too. +(component + (import "host" (instance $h + (export "transform" (func (param "x" s32) (result s32))) + (export "greet" (func (param "name" string) (result string))))) + (alias export $h "transform" (func $transform)) + (core func $transform_core (canon lower (func $transform))) + (core module $m + (import "host" "transform" (func $transform (param i32) (result i32))) + (func (export "run") (param i32) (result i32) + local.get 0 + call $transform) + (func $start + i32.const 21 + call $transform + i32.const 42 + i32.ne + if + unreachable + end) + (start $start) + ) + (core instance $i (instantiate $m + (with "host" (instance (export "transform" (func $transform_core)))))) + (func (export "run") (param "x" s32) (result s32) + (canon lift (core func $i "run"))) +) diff --git a/tests/Components/tiny.wat b/tests/Components/tiny.wat new file mode 100644 index 00000000..ebcb512c --- /dev/null +++ b/tests/Components/tiny.wat @@ -0,0 +1,12 @@ +;; Minimal component with no imports, used to test instantiation and calling. +(component + (core module $m + (func (export "add") (param i32 i32) (result i32) + local.get 0 + local.get 1 + i32.add) + ) + (core instance $i (instantiate $m)) + (func (export "add") (param "a" s32) (param "b" s32) (result s32) + (canon lift (core func $i "add"))) +) diff --git a/tests/Wasmtime.Tests.csproj b/tests/Wasmtime.Tests.csproj index a13c361c..5ffe0d90 100644 --- a/tests/Wasmtime.Tests.csproj +++ b/tests/Wasmtime.Tests.csproj @@ -22,6 +22,7 @@ + From 8872cf6e00141b9555bd1abc3f57dda1eabfbc28 Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Thu, 17 Sep 2026 21:36:54 +1200 Subject: [PATCH 6/9] Add ComponentFunction and ComponentInstance functionality with tests for exported functions and trap behavior --- src/Components/Component.cs | 59 +++++-- src/Components/ComponentExport.cs | 8 +- src/Components/ComponentFunction.cs | 215 +++++++++++++++++++++++++- src/Components/ComponentInstance.cs | 135 +++++++++++++++- tests/ComponentFunctionTests.cs | 229 ++++++++++++++++++++++++++++ tests/Components/trap.wasm | Bin 0 -> 135 bytes tests/Components/trap.wat | 13 ++ 7 files changed, 634 insertions(+), 25 deletions(-) create mode 100644 tests/ComponentFunctionTests.cs create mode 100644 tests/Components/trap.wasm create mode 100644 tests/Components/trap.wat diff --git a/src/Components/Component.cs b/src/Components/Component.cs index 840f1b26..b4b8e126 100644 --- a/src/Components/Component.cs +++ b/src/Components/Component.cs @@ -170,22 +170,59 @@ public static Component DeserializeFile(Engine engine, string path) return new Component(handle); } + /// + /// Looks up an export of this component by name. + /// + /// The name of the export. + /// The export index, or null if there is no such export. + /// Thrown if is null. public ComponentExport? GetExport(string name) { - var ret = Native.wasmtime_component_get_export_index(NativeHandle, null, name, (nuint)name.Length); - if (ret == IntPtr.Zero) - return null; - - return new ComponentExport(ret); + return GetExport(name, null); } - public ComponentExport? GetExport(string name, ComponentExport instance_export_index) + /// + /// Looks up an export of this component by name, within an exported instance. + /// + /// The name of the export. + /// The instance export to look within, or null for the root. + /// The export index, or null if there is no such export. + /// Thrown if is null. + public ComponentExport? GetExport(string name, ComponentExport? instance_export_index) { - var ret = Native.wasmtime_component_get_export_index(NativeHandle, instance_export_index.NativeHandle, name, (nuint)name.Length); - if (ret == IntPtr.Zero) - return null; + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + // The parent is optional, and a null SafeHandle cannot be marshalled, so pass it as a raw pointer. + var parentHandle = instance_export_index?.NativeHandle; + var parentHandleAddedRef = false; + var nameBytes = Encoding.UTF8.GetBytes(name); + + try + { + parentHandle?.DangerousAddRef(ref parentHandleAddedRef); + var parent = parentHandle?.DangerousGetHandle() ?? IntPtr.Zero; - return new ComponentExport(ret); + unsafe + { + fixed (byte* namePtr = nameBytes) + { + var ret = Native.wasmtime_component_get_export_index( + NativeHandle, parent, namePtr, (nuint)nameBytes.Length); + + return ret == IntPtr.Zero ? null : new ComponentExport(ret); + } + } + } + finally + { + if (parentHandleAddedRef) + { + parentHandle!.DangerousRelease(); + } + } } internal class Handle @@ -222,6 +259,6 @@ internal static class Native public static extern IntPtr wasmtime_component_deserialize_file(Engine.Handle engine, string path, out IntPtr handle); [DllImport(Engine.LibraryName)] - public static extern IntPtr wasmtime_component_get_export_index(Handle component, ComponentExport.Handle? instance_export_index, string name, nuint name_len); + public static extern unsafe IntPtr wasmtime_component_get_export_index(Handle component, IntPtr instance_export_index, byte* name, nuint name_len); } } \ No newline at end of file diff --git a/src/Components/ComponentExport.cs b/src/Components/ComponentExport.cs index 4fb03259..ea276d6c 100644 --- a/src/Components/ComponentExport.cs +++ b/src/Components/ComponentExport.cs @@ -1,9 +1,13 @@ -using Microsoft.Win32.SafeHandles; -using System; +using System; using System.Runtime.InteropServices; +using Microsoft.Win32.SafeHandles; namespace Wasmtime.Components; +/// +/// The index of an export within a or , +/// used to look the export up without repeating a name search. +/// public class ComponentExport : IDisposable { diff --git a/src/Components/ComponentFunction.cs b/src/Components/ComponentFunction.cs index 7e390e0d..e60709c8 100644 --- a/src/Components/ComponentFunction.cs +++ b/src/Components/ComponentFunction.cs @@ -1,19 +1,220 @@ -namespace Wasmtime.Components; +using System; +using System.Collections.Generic; +using System.Runtime.InteropServices; + +namespace Wasmtime.Components; /// -/// Represents a Wasmtime function. +/// An exported function of a . /// public class ComponentFunction { - //todo: everything! + /// Size of wasmtime_component_valtype_t: a 1-byte kind padded to 8, plus an 8-byte union. + private const int ValueTypeSize = 16; + + private readonly Store store; + private readonly Native.Func func; + + internal ComponentFunction(Store store, Native.Func func) + { + this.store = store; + this.func = func; + + var type = Native.wasmtime_component_func_type(in func, store.Context.handle); + GC.KeepAlive(store); + + try + { + ParameterCount = (int)Native.wasmtime_component_func_type_param_count(type); + HasResult = ReadHasResult(type); + } + finally + { + Native.wasmtime_component_func_type_delete(type); + } + } + + /// + /// Gets the number of parameters this function takes. + /// + public int ParameterCount { get; } + + /// + /// Gets whether this function returns a value. + /// + /// + /// A component function has at most one result. Note that a WIT function declared to return + /// result<_, E> still has one result here, even though bindings generators + /// usually surface it as returning nothing. + /// + public bool HasResult { get; } + + /// + /// Invokes the function. + /// + /// The arguments, which must match . + /// The result, or null if the function does not return one. + /// Thrown if is null. + /// Thrown if the wrong number of arguments is given. + /// Thrown if the function traps or fails. + public ComponentValue? Call(params ComponentValue[] arguments) + { + return Call((IReadOnlyList)arguments); + } + + /// + /// Invokes the function. + /// + /// The arguments, which must match . + /// The result, or null if the function does not return one. + /// Thrown if is null. + /// Thrown if the wrong number of arguments is given. + /// Thrown if the function traps or fails. + /// + /// A trap leaves the whole unusable for further component calls, which + /// then fail with "cannot enter component instance". This applies to every instance in the + /// store, not just this one, so a store cannot be reused after a trap. + /// + public ComponentValue? Call(IReadOnlyList arguments) + { + if (arguments is null) + { + throw new ArgumentNullException(nameof(arguments)); + } + + var argumentCount = arguments.Count; + if (argumentCount != ParameterCount) + { + throw new ArgumentException( + $"The function takes {ParameterCount} argument(s) but {argumentCount} were given.", + nameof(arguments)); + } + + var resultCount = HasResult ? 1 : 0; + + using (var scope = new ComponentValueMarshaller.AllocationScope()) + { + var argumentBuffer = IntPtr.Zero; + if (argumentCount > 0) + { + argumentBuffer = scope.Allocate(argumentCount * ComponentValueMarshaller.ValueSize); + for (var i = 0; i < argumentCount; i++) + { + if (arguments[i] is null) + { + throw new ArgumentException($"Argument {i} is null.", nameof(arguments)); + } + + ComponentValueMarshaller.Write( + arguments[i], + argumentBuffer + (i * ComponentValueMarshaller.ValueSize), + scope); + } + } + // Wasmtime allocates the contents of the results, so the buffer must start zeroed and + // must only be deleted if the call actually wrote to it. + var resultBuffer = resultCount == 0 + ? IntPtr.Zero + : scope.Allocate(resultCount * ComponentValueMarshaller.ValueSize); + + var error = Native.wasmtime_component_func_call( + in func, + store.Context.handle, + argumentBuffer, + (nuint)argumentCount, + resultBuffer, + (nuint)resultCount); + + GC.KeepAlive(store); + + if (error != IntPtr.Zero) + { + throw WasmtimeException.FromOwnedError(error); + } + + if (resultCount == 0) + { + return null; + } + + try + { + return ComponentValueMarshaller.Read(resultBuffer); + } + finally + { + ComponentValueNative.wasmtime_component_val_delete(resultBuffer); + } + } + } + + private static bool ReadHasResult(IntPtr type) + { + var buffer = Marshal.AllocHGlobal(ValueTypeSize); + try + { + unsafe + { + new Span((void*)buffer, ValueTypeSize).Clear(); + } + + if (!Native.wasmtime_component_func_type_result(type, buffer)) + { + return false; + } + + Native.wasmtime_component_valtype_delete(buffer); + return true; + } + finally + { + Marshal.FreeHGlobal(buffer); + } + } internal static class Native { - // [DllImport(Engine.LibraryName)] - //todo: wasmtime_error_t * wasmtime_component_func_call (const wasmtime_component_func_t *func, wasmtime_context_t *context, const wasmtime_component_val_t *args, size_t args_size, wasmtime_component_val_t *results, size_t results_size) + /// + /// Mirrors wasmtime_component_func_t. The first two fields sit in an anonymous + /// struct, so the trailing field lands at offset 16 and the whole thing is 24 bytes. + /// + [StructLayout(LayoutKind.Explicit, Size = 24)] + internal struct Func + { + [FieldOffset(0)] + public ulong StoreId; + + [FieldOffset(8)] + public uint Private1; + + [FieldOffset(16)] + public uint Private2; + } + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_func_call( + in Func func, + IntPtr context, + IntPtr args, + nuint args_size, + IntPtr results, + nuint results_size); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_func_type(in Func func, IntPtr context); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_func_type_delete(IntPtr ty); + + [DllImport(Engine.LibraryName)] + public static extern nuint wasmtime_component_func_type_param_count(IntPtr ty); + + [DllImport(Engine.LibraryName)] + [return: MarshalAs(UnmanagedType.I1)] + public static extern bool wasmtime_component_func_type_result(IntPtr ty, IntPtr type_ret); - // [DllImport(Engine.LibraryName)] - //todo: wasmtime_error_t * wasmtime_component_func_post_return (const wasmtime_component_func_t *func, wasmtime_context_t *context) + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_valtype_delete(IntPtr ptr); } } \ No newline at end of file diff --git a/src/Components/ComponentInstance.cs b/src/Components/ComponentInstance.cs index 5e991448..37a28d7d 100644 --- a/src/Components/ComponentInstance.cs +++ b/src/Components/ComponentInstance.cs @@ -1,4 +1,6 @@ -using System.Runtime.InteropServices; +using System; +using System.Runtime.InteropServices; +using System.Text; namespace Wasmtime.Components; @@ -21,6 +23,119 @@ internal ComponentInstance(Store store, Native.Instance instance) this.instance = instance; } + /// + /// Looks up an export of this instance by name. + /// + /// The name of the export. + /// The instance export to look within, or null for the root. + /// The export index, or null if there is no such export. + /// Thrown if is null. + public ComponentExport? GetExport(string name, ComponentExport? parent = null) + { + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + var nameBytes = Encoding.UTF8.GetBytes(name); + + // The parent is optional, and a null SafeHandle cannot be marshalled, so pass it as a raw pointer. + var parentSafeHandle = parent?.NativeHandle; + var parentHandleAddedRef = false; + + try + { + parentSafeHandle?.DangerousAddRef(ref parentHandleAddedRef); + var parentHandle = parentSafeHandle?.DangerousGetHandle() ?? IntPtr.Zero; + + unsafe + { + fixed (byte* namePtr = nameBytes) + { + var index = Native.wasmtime_component_instance_get_export_index( + in instance, + store.Context.handle, + parentHandle, + namePtr, + (nuint)nameBytes.Length); + + GC.KeepAlive(store); + + return index == IntPtr.Zero ? null : new ComponentExport(index); + } + } + } + finally + { + if (parentHandleAddedRef) + { + parentSafeHandle!.DangerousRelease(); + } + } + } + + /// + /// Gets an exported function by its export index. + /// + /// The export index of the function. + /// The function, or null if the export is not a function. + /// Thrown if is null. + public ComponentFunction? GetFunction(ComponentExport export) + { + if (export is null) + { + throw new ArgumentNullException(nameof(export)); + } + + var found = Native.wasmtime_component_instance_get_func( + in instance, store.Context.handle, export.NativeHandle, out var func); + + GC.KeepAlive(store); + + return found ? new ComponentFunction(store, func) : null; + } + + /// + /// Gets an exported function from the root of this instance. + /// + /// The name of the function. + /// The function, or null if there is no such exported function. + /// Thrown if is null. + public ComponentFunction? GetFunction(string name) + { + using var export = GetExport(name); + return export is null ? null : GetFunction(export); + } + + /// + /// Gets a function exported by one of this instance's exported interfaces. + /// + /// The name of the interface, such as my:pkg/iface@1.0.0. + /// The name of the function within the interface. + /// The function, or null if there is no such exported function. + /// Thrown if an argument is null. + public ComponentFunction? GetFunction(string interfaceName, string name) + { + if (interfaceName is null) + { + throw new ArgumentNullException(nameof(interfaceName)); + } + + if (name is null) + { + throw new ArgumentNullException(nameof(name)); + } + + using var instanceExport = GetExport(interfaceName); + if (instanceExport is null) + { + return null; + } + + using var export = GetExport(name, instanceExport); + return export is null ? null : GetFunction(export); + } + internal static class Native { /// Mirrors wasmtime_component_instance_t. @@ -34,10 +149,20 @@ internal struct Instance public uint Private; } - //[DllImport(Engine.LibraryName)] - //public static extern IntPtr /* wasmtime_component_export_index_t* */ wasmtime_component_instance_get_export_index (wasmtime_component_instance_t *instance, wasmtime_context_t *context, ComponentExport.Handle instance_export_index, string name, nuint name_len) + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_component_instance_get_export_index( + in Instance instance, + IntPtr context, + IntPtr instance_export_index, + byte* name, + nuint name_len); - // [DllImport(Engine.LibraryName)] - //todo: bool wasmtime_component_instance_get_func (const wasmtime_component_instance_t *instance, wasmtime_context_t *context, const wasmtime_component_export_index_t *export_index, wasmtime_component_func_t *func_out) + [DllImport(Engine.LibraryName)] + [return: MarshalAs(UnmanagedType.I1)] + public static extern bool wasmtime_component_instance_get_func( + in Instance instance, + IntPtr context, + ComponentExport.Handle export_index, + out ComponentFunction.Native.Func func_out); } } \ No newline at end of file diff --git a/tests/ComponentFunctionTests.cs b/tests/ComponentFunctionTests.cs new file mode 100644 index 00000000..f715feee --- /dev/null +++ b/tests/ComponentFunctionTests.cs @@ -0,0 +1,229 @@ +using System; +using System.Collections; +using System.Collections.Generic; +using FluentAssertions; +using Wasmtime.Components; +using Xunit; + +namespace Wasmtime.Tests +{ + public class ComponentFunctionFixture : ComponentFixture { } + + public sealed class ComponentFunctionTests : IClassFixture, IDisposable + { + private readonly ComponentFunctionFixture fixture; + private readonly ComponentLinker linker; + private readonly Store store; + + public ComponentFunctionTests(ComponentFunctionFixture fixture) + { + this.fixture = fixture; + store = new Store(fixture.Engine); + linker = new ComponentLinker(fixture.Engine); + } + + public void Dispose() + { + store.Dispose(); + linker.Dispose(); + } + + private ComponentInstance InstantiateTinyComponent() + { + using var component = fixture.LoadComponent("tiny.wat"); + return linker.Instantiate(store, component); + } + + [Fact] + public void ItCallsAnExportedFunction() + { + var instance = InstantiateTinyComponent(); + + var add = instance.GetFunction("add"); + + add.Should().NotBeNull(); + add!.Call(ComponentValue.S32(2), ComponentValue.S32(3))!.AsS32().Should().Be(5); + } + + [Fact] + public void ItReportsTheFunctionSignature() + { + var instance = InstantiateTinyComponent(); + + var add = instance.GetFunction("add")!; + + add.ParameterCount.Should().Be(2); + add.HasResult.Should().BeTrue(); + } + + [Fact] + public void ItCallsAFunctionRepeatedly() + { + var instance = InstantiateTinyComponent(); + var add = instance.GetFunction("add")!; + + for (var i = 0; i < 100; i++) + { + add.Call(ComponentValue.S32(i), ComponentValue.S32(i))!.AsS32().Should().Be(i * 2); + } + } + + [Theory] + [InlineData(0, 0, 0)] + [InlineData(-1, 1, 0)] + [InlineData(int.MaxValue, 1, int.MinValue)] + [InlineData(int.MinValue, -1, int.MaxValue)] + public void ItRoundTripsArgumentsAndResults(int a, int b, int expected) + { + var instance = InstantiateTinyComponent(); + var add = instance.GetFunction("add")!; + + add.Call(ComponentValue.S32(a), ComponentValue.S32(b))!.AsS32().Should().Be(expected); + } + + [Fact] + public void ItReturnsNullForAnUnknownExport() + { + var instance = InstantiateTinyComponent(); + + instance.GetExport("nope").Should().BeNull(); + instance.GetFunction("nope").Should().BeNull(); + instance.GetFunction("no:such/iface", "nope").Should().BeNull(); + } + + [Fact] + public void ItThrowsForTheWrongNumberOfArguments() + { + var instance = InstantiateTinyComponent(); + var add = instance.GetFunction("add")!; + + add.Invoking(f => f.Call(ComponentValue.S32(1))) + .Should().Throw().WithMessage("*2 argument(s)*1 were given*"); + add.Invoking(f => f.Call()) + .Should().Throw(); + } + + [Fact] + public void ItThrowsForNullArguments() + { + var instance = InstantiateTinyComponent(); + var add = instance.GetFunction("add")!; + + add.Invoking(f => f.Call((ComponentValue[])null!)).Should().Throw(); + instance.Invoking(i => i.GetExport(null!)).Should().Throw(); + instance.Invoking(i => i.GetFunction((ComponentExport)null!)).Should().Throw(); + instance.Invoking(i => i.GetFunction("no:such/iface", null!)).Should().Throw(); + } + + [Fact] + public void ItUsesOneArgumentCountForMarshalling() + { + var instance = InstantiateTinyComponent(); + var add = instance.GetFunction("add")!; + var arguments = new ChangingCountList(ComponentValue.S32(2), ComponentValue.S32(3)); + + add.Call(arguments)!.AsS32().Should().Be(5); + } + + [Fact] + public void ItCallsBackIntoTheHost() + { + var observed = 0; + + using (var root = linker.Root()) + using (var host = root.AddInstance("host")) + { + host.DefineFunction("transform", (arguments, results) => + { + observed = arguments[0].AsS32(); + results[0] = ComponentValue.S32(observed * 2); + }); + + host.DefineFunction("greet", (arguments, results) => + results[0] = ComponentValue.String($"hello, {arguments[0].AsString()}")); + } + + using var component = fixture.LoadComponent("host-import.wat"); + var instance = linker.Instantiate(store, component); + + instance.GetFunction("run")!.Call(ComponentValue.S32(21))!.AsS32().Should().Be(42); + observed.Should().Be(21); + } + + [Fact] + public void ItSurfacesHostExceptionsAsTraps() + { + using (var root = linker.Root()) + using (var host = root.AddInstance("host")) + { + host.DefineFunction("transform", (arguments, results) => + throw new InvalidOperationException("host went bang")); + host.DefineFunction("greet", (arguments, results) => + results[0] = ComponentValue.String(string.Empty)); + } + + using var component = fixture.LoadComponent("host-import.wat"); + linker.Invoking(l => l.Instantiate(store, component)) + .Should().Throw() + .WithMessage("*host went bang*"); + } + + [Fact] + public void ItThrowsWhenAHostFunctionDoesNotSetItsResult() + { + using (var root = linker.Root()) + using (var host = root.AddInstance("host")) + { + host.DefineFunction("transform", (arguments, results) => { }); + host.DefineFunction("greet", (arguments, results) => + results[0] = ComponentValue.String(string.Empty)); + } + + using var component = fixture.LoadComponent("host-import.wat"); + linker.Invoking(l => l.Instantiate(store, component)) + .Should().Throw() + .WithMessage("*did not assign result*"); + } + + /// + /// A trap poisons its entire Store, not just the instance that trapped. This holds for + /// guest-side traps and for host callbacks that throw, so it is pinned here: sharing a + /// Store across calls that may trap silently breaks every later call. + /// + [Fact] + public void ATrapPoisonsTheWholeStore() + { + using var tiny = fixture.LoadComponent("tiny.wat"); + using var trap = fixture.LoadComponent("trap.wat"); + + var add = linker.Instantiate(store, tiny).GetFunction("add")!; + add.Call(ComponentValue.S32(1), ComponentValue.S32(1))!.AsS32().Should().Be(2); + + linker.Instantiate(store, trap).GetFunction("boom")! + .Invoking(f => f.Call()) + .Should().Throw(); + + add.Invoking(f => f.Call(ComponentValue.S32(1), ComponentValue.S32(1))) + .Should().Throw() + .WithMessage("*cannot enter component instance*"); + } + private sealed class ChangingCountList : IReadOnlyList + { + private readonly ComponentValue[] items; + private int countReads; + + public ChangingCountList(params ComponentValue[] items) + { + this.items = items; + } + + public int Count => countReads++ == 0 ? items.Length : items.Length + 1; + + public ComponentValue this[int index] => items[index]; + + public IEnumerator GetEnumerator() => ((IEnumerable)items).GetEnumerator(); + + IEnumerator IEnumerable.GetEnumerator() => items.GetEnumerator(); + } + } +} diff --git a/tests/Components/trap.wasm b/tests/Components/trap.wasm new file mode 100644 index 0000000000000000000000000000000000000000..695ea598a6e6988ba82dfb5d07799ab683bdafa2 GIT binary patch literal 135 zcmW-aI}XAy5JcZNwiALtf(vj0lq=*Y1eAheZ6v5@xH{`p^Vz3ambV=cns6sTNjLQC zg6Idud#wgzQU`l>u`IR{WFa=VPnzMIM-O6yhujVhd$$!WUXG7yuUgjfxwZvmPX5$f H&#uourrr|m literal 0 HcmV?d00001 diff --git a/tests/Components/trap.wat b/tests/Components/trap.wat new file mode 100644 index 00000000..958d79d6 --- /dev/null +++ b/tests/Components/trap.wat @@ -0,0 +1,13 @@ +;; Component whose exported function always traps, used to test trap behaviour. +;; +;; Rebuild with: +;; wasm-tools parse tests/Components/trap.wat -o tests/Components/trap.wasm +(component + (core module $m + (func (export "boom") (result i32) + unreachable) + ) + (core instance $i (instantiate $m)) + (func (export "boom") (result s32) + (canon lift (core func $i "boom"))) +) From d6303fe53b6088b884c9d0a7b4a83e68f0e5fdf2 Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Thu, 17 Sep 2026 21:49:26 +1200 Subject: [PATCH 7/9] Add component model support and examples - Updated README.md to include documentation for the component model in Wasmtime. - Added a new project for examples demonstrating the component model. - Implemented a sample component in WebAssembly text format (WAT) that imports a host function and exports a greeting function. - Introduced methods in the Component class to create components from WAT text and files. - Added tests for component functionality, including string handling and list results. - Updated solution file to include the new examples and component projects. --- .github/workflows/main.yml | 2 +- Examples.sln | 6 ++ examples/component/Program.cs | 27 +++++++ examples/component/component.csproj | 14 ++++ examples/component/greeter.wat | 86 ++++++++++++++++++++++ src/Components/Component.cs | 21 ++++++ tests/ComponentFixture.cs | 4 ++ tests/ComponentFunctionTests.cs | 69 ++++++++++++++++++ tests/ComponentTests.cs | 107 ++++++++++++++++++++++++++++ tests/Components/strings.wat | 75 +++++++++++++++++++ tests/Components/trap.wat | 3 - 11 files changed, 410 insertions(+), 4 deletions(-) create mode 100644 examples/component/Program.cs create mode 100644 examples/component/component.csproj create mode 100644 examples/component/greeter.wat create mode 100644 tests/ComponentTests.cs create mode 100644 tests/Components/strings.wat diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml index 3ded1423..1eb12d4b 100644 --- a/.github/workflows/main.yml +++ b/.github/workflows/main.yml @@ -55,7 +55,7 @@ jobs: - name: Run examples shell: bash env: - EXAMPLES: externref funcref global hello memory table consumefuel + EXAMPLES: externref funcref global hello memory table consumefuel component run: | for e in $EXAMPLES; do cd examples/$e && dotnet run -c ${{ matrix.config }} && cd ../..; done - name: Create package diff --git a/Examples.sln b/Examples.sln index b0d31c21..9867cb99 100644 --- a/Examples.sln +++ b/Examples.sln @@ -21,6 +21,8 @@ Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "consumefuel", "examples\con EndProject Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "storedata", "examples\storedata\storedata.csproj", "{E8749BAF-9D8B-4CF9-ACFF-490E86D52A20}" EndProject +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "component", "examples\component\component.csproj", "{35CE1A43-6F41-4747-82DB-6BFEE10297EC}" +EndProject Global GlobalSection(SolutionConfigurationPlatforms) = preSolution Debug|Any CPU = Debug|Any CPU @@ -63,6 +65,10 @@ Global {E8749BAF-9D8B-4CF9-ACFF-490E86D52A20}.Debug|Any CPU.Build.0 = Debug|Any CPU {E8749BAF-9D8B-4CF9-ACFF-490E86D52A20}.Release|Any CPU.ActiveCfg = Release|Any CPU {E8749BAF-9D8B-4CF9-ACFF-490E86D52A20}.Release|Any CPU.Build.0 = Release|Any CPU + {35CE1A43-6F41-4747-82DB-6BFEE10297EC}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {35CE1A43-6F41-4747-82DB-6BFEE10297EC}.Debug|Any CPU.Build.0 = Debug|Any CPU + {35CE1A43-6F41-4747-82DB-6BFEE10297EC}.Release|Any CPU.ActiveCfg = Release|Any CPU + {35CE1A43-6F41-4747-82DB-6BFEE10297EC}.Release|Any CPU.Build.0 = Release|Any CPU EndGlobalSection GlobalSection(SolutionProperties) = preSolution HideSolutionNode = FALSE diff --git a/examples/component/Program.cs b/examples/component/Program.cs new file mode 100644 index 00000000..db7beb2e --- /dev/null +++ b/examples/component/Program.cs @@ -0,0 +1,27 @@ +using System; +using System.IO; +using Wasmtime; +using Wasmtime.Components; + +using var engine = new Engine(new Config().WithComponentModel(true)); +using var component = Component.FromTextFile(engine, Path.Combine(AppContext.BaseDirectory, "greeter.wat")); +using var linker = new ComponentLinker(engine); +using var store = new Store(engine); + +using (var root = linker.Root()) +using (var host = root.AddInstance("host")) +{ + host.DefineFunction("name", (arguments, results) => + results[0] = ComponentValue.String("WebAssembly")); +} + +var instance = linker.Instantiate(store, component); + +var greet = instance.GetFunction("greet"); +if (greet is null) +{ + Console.WriteLine("error: greet export is missing"); + return; +} + +Console.WriteLine(greet.Call(ComponentValue.String("Wasmtime"))!.AsString()); diff --git a/examples/component/component.csproj b/examples/component/component.csproj new file mode 100644 index 00000000..8bbbbb19 --- /dev/null +++ b/examples/component/component.csproj @@ -0,0 +1,14 @@ + + + + Exe + net9.0 + enable + + + + + + + + diff --git a/examples/component/greeter.wat b/examples/component/greeter.wat new file mode 100644 index 00000000..2f613b44 --- /dev/null +++ b/examples/component/greeter.wat @@ -0,0 +1,86 @@ +;; A component that combines a caller-provided string with a host-provided name. +;; +;; import name: func() -> string +;; export greet: func(string) -> string +;; +;; The core module needs an exported memory and realloc so that strings can cross +;; the component boundary. +(component + (import "host" (instance $h + (export "name" (func (result string))))) + (alias export $h "name" (func $name)) + + (core module $libc + (memory (export "memory") 1) + (global $bump (mut i32) (i32.const 1024)) + (func (export "realloc") (param $old i32) (param $oldsz i32) (param $align i32) (param $n i32) (result i32) + (local $r i32) + (local.set $r (global.get $bump)) + (global.set $bump + (i32.and (i32.add (i32.add (global.get $bump) (local.get $n)) (i32.const 7)) (i32.const -8))) + (local.get $r)) + ) + (core instance $libc_i (instantiate $libc)) + + (core func $name_core + (canon lower (func $name) + (memory $libc_i "memory") + (realloc (func $libc_i "realloc")) + string-encoding=utf8)) + + (core module $m + (import "libc" "memory" (memory 1)) + (import "libc" "realloc" (func $realloc (param i32 i32 i32 i32) (result i32))) + (import "host" "name" (func $name (param i32))) + + (func (export "greet") (param $ptr i32) (param $len i32) (result i32) + (local $host_result i32) + (local $host_ptr i32) + (local $host_len i32) + (local $out i32) + (local $suffix i32) + (local $ret i32) + + ;; Ask the host for its name, then build "hello, " ++ name ++ " from " ++ host name. + (local.set $host_result (call $realloc (i32.const 0) (i32.const 0) (i32.const 4) (i32.const 8))) + (call $name (local.get $host_result)) + (local.set $host_ptr (i32.load (local.get $host_result))) + (local.set $host_len (i32.load offset=4 (local.get $host_result))) + + (local.set $out + (call $realloc (i32.const 0) (i32.const 0) (i32.const 1) + (i32.add (i32.const 13) (i32.add (local.get $len) (local.get $host_len))))) + (i32.store8 (local.get $out) (i32.const 104)) + (i32.store8 offset=1 (local.get $out) (i32.const 101)) + (i32.store8 offset=2 (local.get $out) (i32.const 108)) + (i32.store8 offset=3 (local.get $out) (i32.const 108)) + (i32.store8 offset=4 (local.get $out) (i32.const 111)) + (i32.store8 offset=5 (local.get $out) (i32.const 44)) + (i32.store8 offset=6 (local.get $out) (i32.const 32)) + (memory.copy (i32.add (local.get $out) (i32.const 7)) (local.get $ptr) (local.get $len)) + + (local.set $suffix (i32.add (i32.add (local.get $out) (i32.const 7)) (local.get $len))) + (i32.store8 (local.get $suffix) (i32.const 32)) + (i32.store8 offset=1 (local.get $suffix) (i32.const 102)) + (i32.store8 offset=2 (local.get $suffix) (i32.const 114)) + (i32.store8 offset=3 (local.get $suffix) (i32.const 111)) + (i32.store8 offset=4 (local.get $suffix) (i32.const 109)) + (i32.store8 offset=5 (local.get $suffix) (i32.const 32)) + (memory.copy (i32.add (local.get $suffix) (i32.const 6)) (local.get $host_ptr) (local.get $host_len)) + + (local.set $ret (call $realloc (i32.const 0) (i32.const 0) (i32.const 4) (i32.const 8))) + (i32.store (local.get $ret) (local.get $out)) + (i32.store offset=4 (local.get $ret) + (i32.add (i32.const 13) (i32.add (local.get $len) (local.get $host_len)))) + (local.get $ret)) + ) + (core instance $i (instantiate $m + (with "libc" (instance $libc_i)) + (with "host" (instance (export "name" (func $name_core)))))) + + (func (export "greet") (param "name" string) (result string) + (canon lift (core func $i "greet") + (memory $libc_i "memory") + (realloc (func $libc_i "realloc")) + string-encoding=utf8)) +) diff --git a/src/Components/Component.cs b/src/Components/Component.cs index b4b8e126..1a5b745a 100644 --- a/src/Components/Component.cs +++ b/src/Components/Component.cs @@ -1,4 +1,5 @@ using System; +using System.IO; using System.Runtime.InteropServices; using System.Text; using Microsoft.Win32.SafeHandles; @@ -102,6 +103,23 @@ public static Component FromText(Engine engine, string text) } } + /// + /// Creates a from a file in the WebAssembly text format. + /// + /// The engine to use for the component. + /// The path to the file. + /// Returns a new . + /// Thrown if an argument is null. + public static Component FromTextFile(Engine engine, string path) + { + if (path is null) + { + throw new ArgumentNullException(nameof(path)); + } + + return FromText(engine, File.ReadAllText(path)); + } + /// /// This function serializes compiled component artifacts as blob data. /// @@ -260,5 +278,8 @@ internal static class Native [DllImport(Engine.LibraryName)] public static extern unsafe IntPtr wasmtime_component_get_export_index(Handle component, IntPtr instance_export_index, byte* name, nuint name_len); + + [DllImport(Engine.LibraryName)] + public static extern unsafe IntPtr wasmtime_wat2wasm(byte* text, nuint len, out ByteArray bytes); } } \ No newline at end of file diff --git a/tests/ComponentFixture.cs b/tests/ComponentFixture.cs index 9936b7fc..a9734372 100644 --- a/tests/ComponentFixture.cs +++ b/tests/ComponentFixture.cs @@ -22,6 +22,10 @@ public virtual Config GetEngineConfig() public Component LoadComponent(string fileName) => Component.FromText(Engine, File.ReadAllText(Path.Combine("Components", fileName))); + public Store CreateStore() => new Store(Engine); + + public Component Strings() => LoadComponent("strings.wat"); + public void Dispose() { if (!(Store is null)) diff --git a/tests/ComponentFunctionTests.cs b/tests/ComponentFunctionTests.cs index f715feee..195d4d61 100644 --- a/tests/ComponentFunctionTests.cs +++ b/tests/ComponentFunctionTests.cs @@ -224,6 +224,75 @@ public ChangingCountList(params ComponentValue[] items) public IEnumerator GetEnumerator() => ((IEnumerable)items).GetEnumerator(); IEnumerator IEnumerable.GetEnumerator() => items.GetEnumerator(); + } + + private ComponentInstance InstantiateStrings(ComponentLinker linker, Store store) + { + using var component = fixture.Strings(); + return linker.Instantiate(store, component); + } + + [Theory] + [InlineData("world", "hi world")] + [InlineData("", "hi ")] + [InlineData("h\u00e9llo", "hi h\u00e9llo")] + [InlineData("\u4e16\u754c", "hi \u4e16\u754c")] + [InlineData("\ud83c\udf89", "hi \ud83c\udf89")] + public void ItPassesAndReturnsStrings(string name, string expected) + { + using var store = fixture.CreateStore(); + using var linker = new ComponentLinker(fixture.Engine); + var greet = InstantiateStrings(linker, store).GetFunction("greet")!; + + greet.Call(ComponentValue.String(name))!.AsString().Should().Be(expected); + } + + /// + /// The guest reports the byte length it was given, so this fails if the argument was + /// encoded as anything other than UTF-8. + /// + [Theory] + [InlineData("abc", 3)] + [InlineData("h\u00e9llo", 6)] + [InlineData("\u4e16\u754c", 6)] + [InlineData("\ud83c\udf89", 4)] + [InlineData("", 0)] + public void ItPassesStringsAsUtf8(string name, int expectedByteLength) + { + using var store = fixture.CreateStore(); + using var linker = new ComponentLinker(fixture.Engine); + var length = InstantiateStrings(linker, store).GetFunction("length")!; + + length.Call(ComponentValue.String(name))!.AsS32().Should().Be(expectedByteLength); + } + + [Fact] + public void ItReturnsALongStringFromTheGuest() + { + using var store = fixture.CreateStore(); + using var linker = new ComponentLinker(fixture.Engine); + var greet = InstantiateStrings(linker, store).GetFunction("greet")!; + + var name = new string('x', 10_000); + + greet.Call(ComponentValue.String(name))!.AsString().Should().Be("hi " + name); + } + + [Fact] + public void ItReturnsAListFromTheGuest() + { + using var store = fixture.CreateStore(); + using var linker = new ComponentLinker(fixture.Engine); + var numbers = InstantiateStrings(linker, store).GetFunction("numbers")!; + + numbers.ParameterCount.Should().Be(0); + + var result = numbers.Call()!.AsList(); + + result.Should().HaveCount(3); + result[0].AsS32().Should().Be(10); + result[1].AsS32().Should().Be(20); + result[2].AsS32().Should().Be(30); } } } diff --git a/tests/ComponentTests.cs b/tests/ComponentTests.cs new file mode 100644 index 00000000..85100a89 --- /dev/null +++ b/tests/ComponentTests.cs @@ -0,0 +1,107 @@ +using System; +using System.IO; +using FluentAssertions; +using Wasmtime.Components; +using Xunit; + +namespace Wasmtime.Tests +{ + public class ComponentTestsFixture : ComponentFixture { } + + public sealed class ComponentTests : IClassFixture + { + private const string Text = @" +(component + (core module $m + (func (export ""add"") (param i32 i32) (result i32) + local.get 0 + local.get 1 + i32.add) + ) + (core instance $i (instantiate $m)) + (func (export ""add"") (param ""a"" s32) (param ""b"" s32) (result s32) + (canon lift (core func $i ""add""))) +)"; + + private readonly ComponentTestsFixture fixture; + + public ComponentTests(ComponentTestsFixture fixture) + { + this.fixture = fixture; + } + + [Fact] + public void ItLoadsAComponentFromText() + { + using var store = fixture.CreateStore(); + using var component = Component.FromText(fixture.Engine, Text); + using var linker = new ComponentLinker(fixture.Engine); + + var add = linker.Instantiate(store, component).GetFunction("add")!; + + add.Call(ComponentValue.S32(40), ComponentValue.S32(2))!.AsS32().Should().Be(42); + } + + [Fact] + public void ItLoadsAComponentFromATextFile() + { + using var store = fixture.CreateStore(); + using var component = Component.FromTextFile(fixture.Engine, "Components/tiny.wat"); + using var linker = new ComponentLinker(fixture.Engine); + + var add = linker.Instantiate(store, component).GetFunction("add")!; + + add.Call(ComponentValue.S32(1), ComponentValue.S32(2))!.AsS32().Should().Be(3); + } + + [Fact] + public void ItThrowsForInvalidText() + { + FluentActions.Invoking(() => Component.FromText(fixture.Engine, "(component (this is not valid)")) + .Should().Throw(); + } + + [Fact] + public void ItThrowsForInvalidBytes() + { + FluentActions.Invoking(() => Component.FromBytes(fixture.Engine, new byte[] { 1, 2, 3, 4 })) + .Should().Throw(); + } + + [Fact] + public void ItThrowsForNullArguments() + { + FluentActions.Invoking(() => Component.FromText(null!, Text)).Should().Throw(); + FluentActions.Invoking(() => Component.FromText(fixture.Engine, null!)).Should().Throw(); + FluentActions.Invoking(() => Component.FromTextFile(fixture.Engine, null!)).Should().Throw(); + } + + [Fact] + public void ItRoundTripsThroughSerialization() + { + using var store = fixture.CreateStore(); + using var original = Component.FromText(fixture.Engine, Text); + + var serialized = original.Serialize(); + serialized.Should().NotBeEmpty(); + + using var restored = Component.Deserialize(fixture.Engine, serialized); + using var linker = new ComponentLinker(fixture.Engine); + + var add = linker.Instantiate(store, restored).GetFunction("add")!; + + add.Call(ComponentValue.S32(2), ComponentValue.S32(2))!.AsS32().Should().Be(4); + } + + [Fact] + public void ItLooksUpExportsOnTheComponentItself() + { + using var component = Component.FromText(fixture.Engine, Text); + + using var export = component.GetExport("add"); + + export.Should().NotBeNull(); + component.GetExport("nope").Should().BeNull(); + } + } +} diff --git a/tests/Components/strings.wat b/tests/Components/strings.wat new file mode 100644 index 00000000..28a2d851 --- /dev/null +++ b/tests/Components/strings.wat @@ -0,0 +1,75 @@ +;; Component exercising heap-carrying values across the guest boundary: a string parameter, +;; a string result and a list result. Lifting these requires the core module to export a +;; memory and a realloc, and to return results indirectly through a return area. +;; +;; The allocator is a bump allocator that never frees, which is all these tests need. +(component + (core module $m + (memory (export "memory") 1) + (global $bump (mut i32) (i32.const 1024)) + + (func $alloc (param $n i32) (result i32) + (local $r i32) + (local.set $r (global.get $bump)) + (global.set $bump + (i32.and + (i32.add (i32.add (global.get $bump) (local.get $n)) (i32.const 7)) + (i32.const -8))) + (local.get $r)) + + (func (export "realloc") (param $old i32) (param $oldsz i32) (param $align i32) (param $newsz i32) (result i32) + (call $alloc (local.get $newsz))) + + ;; greet(name) -> "hi " ++ name + (func (export "greet") (param $ptr i32) (param $len i32) (result i32) + (local $out i32) + (local $ret i32) + (local.set $out (call $alloc (i32.add (i32.const 3) (local.get $len)))) + (i32.store8 (local.get $out) (i32.const 104)) + (i32.store8 offset=1 (local.get $out) (i32.const 105)) + (i32.store8 offset=2 (local.get $out) (i32.const 32)) + (memory.copy + (i32.add (local.get $out) (i32.const 3)) + (local.get $ptr) + (local.get $len)) + (local.set $ret (call $alloc (i32.const 8))) + (i32.store (local.get $ret) (local.get $out)) + (i32.store offset=4 (local.get $ret) (i32.add (i32.const 3) (local.get $len))) + (local.get $ret)) + + ;; length(name) -> byte length of the string as seen by the guest + (func (export "length") (param $ptr i32) (param $len i32) (result i32) + (local.get $len)) + + ;; numbers() -> [10, 20, 30] + (func (export "numbers") (result i32) + (local $data i32) + (local $ret i32) + (local.set $data (call $alloc (i32.const 12))) + (i32.store (local.get $data) (i32.const 10)) + (i32.store offset=4 (local.get $data) (i32.const 20)) + (i32.store offset=8 (local.get $data) (i32.const 30)) + (local.set $ret (call $alloc (i32.const 8))) + (i32.store (local.get $ret) (local.get $data)) + (i32.store offset=4 (local.get $ret) (i32.const 3)) + (local.get $ret)) + ) + (core instance $i (instantiate $m)) + + (func (export "greet") (param "name" string) (result string) + (canon lift (core func $i "greet") + (memory $i "memory") + (realloc (func $i "realloc")) + string-encoding=utf8)) + + (func (export "length") (param "name" string) (result s32) + (canon lift (core func $i "length") + (memory $i "memory") + (realloc (func $i "realloc")) + string-encoding=utf8)) + + (func (export "numbers") (result (list s32)) + (canon lift (core func $i "numbers") + (memory $i "memory") + (realloc (func $i "realloc")))) +) diff --git a/tests/Components/trap.wat b/tests/Components/trap.wat index 958d79d6..a2c9d271 100644 --- a/tests/Components/trap.wat +++ b/tests/Components/trap.wat @@ -1,7 +1,4 @@ ;; Component whose exported function always traps, used to test trap behaviour. -;; -;; Rebuild with: -;; wasm-tools parse tests/Components/trap.wat -o tests/Components/trap.wasm (component (core module $m (func (export "boom") (result i32) From 0b9ad096c54e378b9b80fe734936cd57cc1322ed Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Sat, 3 Oct 2026 15:53:28 +1300 Subject: [PATCH 8/9] Optimize memory management in ComponentLinkerInstance and ComponentValueMarshaller --- src/Components/ComponentLinkerInstance.cs | 27 +++++++++++++++++++--- src/Components/ComponentValueMarshaller.cs | 20 ++++++++-------- 2 files changed, 34 insertions(+), 13 deletions(-) diff --git a/src/Components/ComponentLinkerInstance.cs b/src/Components/ComponentLinkerInstance.cs index 85df42ef..3b4d6092 100644 --- a/src/Components/ComponentLinkerInstance.cs +++ b/src/Components/ComponentLinkerInstance.cs @@ -1,4 +1,5 @@ using System; +using System.Buffers; using System.Runtime.InteropServices; using System.Text; using Microsoft.Win32.SafeHandles; @@ -179,16 +180,24 @@ private static IntPtr Invoke( IntPtr results, int resultCount) { + ComponentValue[]? arguments = null; + ComponentValue[]? produced = null; + try { - var arguments = argumentCount == 0 ? Array.Empty() : new ComponentValue[argumentCount]; + arguments = argumentCount == 0 + ? Array.Empty() + : ArrayPool.Shared.Rent(argumentCount); for (var i = 0; i < argumentCount; i++) { arguments[i] = ComponentValueMarshaller.Read(args + (i * ComponentValueMarshaller.ValueSize)); } - var produced = resultCount == 0 ? Array.Empty() : new ComponentValue[resultCount]; - callback(arguments, produced); + produced = resultCount == 0 + ? Array.Empty() + : ArrayPool.Shared.Rent(resultCount); + Array.Clear(produced, 0, resultCount); + callback(arguments.AsSpan(0, argumentCount), produced.AsSpan(0, resultCount)); for (var i = 0; i < resultCount; i++) { @@ -209,6 +218,18 @@ private static IntPtr Invoke( { return HandleCallbackException(ex); } + finally + { + if (arguments is { Length: > 0 }) + { + ArrayPool.Shared.Return(arguments, clearArray: true); + } + + if (produced is { Length: > 0 }) + { + ArrayPool.Shared.Return(produced, clearArray: true); + } + } } /// diff --git a/src/Components/ComponentValueMarshaller.cs b/src/Components/ComponentValueMarshaller.cs index 98d2c66d..9dbaa307 100644 --- a/src/Components/ComponentValueMarshaller.cs +++ b/src/Components/ComponentValueMarshaller.cs @@ -50,11 +50,6 @@ public sealed class AllocationScope : IDisposable private readonly List allocations = new List(); private bool disposed; - ~AllocationScope() - { - FreeAllocations(); - } - /// /// Allocates zeroed native memory whose lifetime is bound to this scope. /// @@ -88,7 +83,6 @@ public void Dispose() disposed = true; FreeAllocations(); - GC.SuppressFinalize(this); } private void FreeAllocations() @@ -284,13 +278,19 @@ private static IntPtr WriteBoxed(ComponentValue value, AllocationScope scope) private static void WriteName(string text, IntPtr destination, AllocationScope scope) { - var bytes = Encoding.UTF8.GetBytes(text); + var byteCount = Encoding.UTF8.GetByteCount(text); // A zero-length allocation would still need a non-null pointer, so always take at least one byte. - var buffer = scope.Allocate(Math.Max(bytes.Length, 1)); - Marshal.Copy(bytes, 0, buffer, bytes.Length); + var buffer = scope.Allocate(Math.Max(byteCount, 1)); + unsafe + { + fixed (char* textPtr = text) + { + Encoding.UTF8.GetBytes(textPtr, text.Length, (byte*)buffer, byteCount); + } + } - Marshal.WriteIntPtr(destination, (IntPtr)bytes.Length); + Marshal.WriteIntPtr(destination, (IntPtr)byteCount); Marshal.WriteIntPtr(destination + VectorDataOffset, buffer); } From 3afc10cf6981e9a31ff53a305ce8aabf6d7ceb03 Mon Sep 17 00:00:00 2001 From: Phillip Cao Date: Sat, 3 Oct 2026 16:19:25 +1300 Subject: [PATCH 9/9] Add native interop methods and enhance ComponentValueMarshaller for better memory management --- src/Components/ComponentValue.cs | 15 ++ src/Components/ComponentValueMarshaller.cs | 239 ++++++++++++++++++--- tests/ComponentValueLayoutTests.cs | 15 ++ 3 files changed, 243 insertions(+), 26 deletions(-) diff --git a/src/Components/ComponentValue.cs b/src/Components/ComponentValue.cs index 3b4eee7f..ed4e1bb9 100644 --- a/src/Components/ComponentValue.cs +++ b/src/Components/ComponentValue.cs @@ -465,6 +465,21 @@ public override string ToString() internal static class ComponentValueNative { + [DllImport(Engine.LibraryName)] + public static extern void wasm_byte_vec_new_uninitialized(IntPtr value, nuint size); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_vallist_new_uninit(IntPtr value, nuint size); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_valrecord_new_uninit(IntPtr value, nuint size); + + [DllImport(Engine.LibraryName)] + public static extern void wasmtime_component_valtuple_new_uninit(IntPtr value, nuint size); + + [DllImport(Engine.LibraryName)] + public static extern IntPtr wasmtime_component_val_new(IntPtr value); + /// /// Performs a deep copy of into . The contents /// of are owned by Wasmtime and must be released with diff --git a/src/Components/ComponentValueMarshaller.cs b/src/Components/ComponentValueMarshaller.cs index 9dbaa307..c132928b 100644 --- a/src/Components/ComponentValueMarshaller.cs +++ b/src/Components/ComponentValueMarshaller.cs @@ -19,9 +19,8 @@ namespace Wasmtime.Components; /// corrupts the heap. /// /// -/// The native structures are written field by field rather than through the -/// wasmtime_component_vallist_new family of helpers, because a value must be built in -/// place inside an array element or a record entry. +/// Argument values are written into caller-owned storage. Host results use Wasmtime's own +/// constructors so their nested allocations have the correct owner. /// /// internal static class ComponentValueMarshaller @@ -104,9 +103,79 @@ private void FreeAllocations() /// The scope owning any secondary allocations the value needs. public static void Write(ComponentValue value, IntPtr destination, AllocationScope scope) { + if (WritePrimitive(value, destination)) + { + return; + } + Marshal.WriteByte(destination, (byte)value.Kind); var payload = destination + ValuePayloadOffset; + switch (value.Kind) + { + case ComponentValueKind.String: + case ComponentValueKind.Enum: + WriteName(value.Text ?? string.Empty, payload, scope); + break; + + case ComponentValueKind.List: + case ComponentValueKind.Tuple: + WriteVector(value.Items, payload, scope); + break; + + case ComponentValueKind.Record: + WriteRecord(value.Fields, payload, scope); + break; + + case ComponentValueKind.Option: + Marshal.WriteIntPtr(payload, value.Flag ? WriteBoxed(value.Payload!, scope) : IntPtr.Zero); + break; + + case ComponentValueKind.Result: + Marshal.WriteByte(payload, (byte)(value.Flag ? 1 : 0)); + Marshal.WriteIntPtr( + payload + VectorDataOffset, + value.Payload is null ? IntPtr.Zero : WriteBoxed(value.Payload, scope)); + break; + + default: + throw new NotSupportedException( + $"Writing component values of kind {value.Kind} is not supported."); + } + } + + /// + /// Writes a value into storage whose contents Wasmtime will free, as required for the + /// results of a host-defined function. + /// + /// The value to write. + /// A pointer to bytes of storage. + /// + /// Heap-backed values are constructed with Wasmtime's allocator so Wasmtime can safely + /// deallocate the result after the host callback returns. + /// + public static void WriteOwned(ComponentValue value, IntPtr destination) + { + ClearValue(destination); + try + { + if (!WritePrimitive(value, destination)) + { + WriteOwnedHeapValue(value, destination); + } + } + catch + { + ComponentValueNative.wasmtime_component_val_delete(destination); + ClearValue(destination); + throw; + } + } + + private static bool WritePrimitive(ComponentValue value, IntPtr destination) + { + var payload = destination + ValuePayloadOffset; + switch (value.Kind) { case ComponentValueKind.Bool: @@ -141,29 +210,52 @@ public static void Write(ComponentValue value, IntPtr destination, AllocationSco Marshal.WriteInt64(payload, BitConverter.DoubleToInt64Bits(value.Real)); break; + default: + return false; + } + + Marshal.WriteByte(destination, (byte)value.Kind); + return true; + } + + private static void WriteOwnedHeapValue(ComponentValue value, IntPtr destination) + { + var payload = destination + ValuePayloadOffset; + + switch (value.Kind) + { case ComponentValueKind.String: case ComponentValueKind.Enum: - WriteName(value.Text ?? string.Empty, payload, scope); + Marshal.WriteByte(destination, (byte)value.Kind); + WriteOwnedName(value.Text ?? string.Empty, payload); break; case ComponentValueKind.List: case ComponentValueKind.Tuple: - WriteVector(value.Items, payload, scope); + WriteOwnedVector(value.Items, destination, value.Kind); break; case ComponentValueKind.Record: - WriteRecord(value.Fields, payload, scope); + WriteOwnedRecord(value.Fields, destination); break; case ComponentValueKind.Option: - Marshal.WriteIntPtr(payload, value.Flag ? WriteBoxed(value.Payload!, scope) : IntPtr.Zero); + Marshal.WriteByte(destination, (byte)value.Kind); + if (value.Flag) + { + Marshal.WriteIntPtr(payload, WriteOwnedBoxed(value.Payload!)); + } + break; case ComponentValueKind.Result: + Marshal.WriteByte(destination, (byte)value.Kind); Marshal.WriteByte(payload, (byte)(value.Flag ? 1 : 0)); - Marshal.WriteIntPtr( - payload + VectorDataOffset, - value.Payload is null ? IntPtr.Zero : WriteBoxed(value.Payload, scope)); + if (value.Payload is not null) + { + Marshal.WriteIntPtr(payload + VectorDataOffset, WriteOwnedBoxed(value.Payload)); + } + break; default: @@ -172,24 +264,119 @@ public static void Write(ComponentValue value, IntPtr destination, AllocationSco } } - /// - /// Writes a value into storage whose contents Wasmtime will free, as required for the - /// results of a host-defined function. - /// - /// The value to write. - /// A pointer to bytes of storage. - /// - /// The value is first built with our own allocator and then deep-copied by Wasmtime, so that - /// every heap payload it ends up owning came from its own allocator. The extra copy is the - /// price of not having to mirror each vec constructor. - /// - public static void WriteOwned(ComponentValue value, IntPtr destination) + private static void WriteOwnedName(string text, IntPtr destination) + { + var byteCount = Encoding.UTF8.GetByteCount(text); + ComponentValueNative.wasm_byte_vec_new_uninitialized(destination, (nuint)byteCount); + + var data = Marshal.ReadIntPtr(destination + VectorDataOffset); + unsafe + { + fixed (char* textPtr = text) + { + Encoding.UTF8.GetBytes(textPtr, text.Length, (byte*)data, byteCount); + } + } + } + + private static void WriteOwnedVector( + IReadOnlyList items, + IntPtr destination, + ComponentValueKind kind) + { + var count = items.Count; + checked + { + _ = count * ValueSize; + } + + Marshal.WriteByte(destination, (byte)kind); + var vector = destination + ValuePayloadOffset; + if (kind == ComponentValueKind.List) + { + ComponentValueNative.wasmtime_component_vallist_new_uninit(vector, (nuint)count); + } + else + { + ComponentValueNative.wasmtime_component_valtuple_new_uninit(vector, (nuint)count); + } + + var element = Marshal.ReadIntPtr(vector + VectorDataOffset); + for (var i = 0; i < count; i++) + { + ClearValue(element + (i * ValueSize)); + } + + for (var i = 0; i < count; i++) + { + WriteOwned(items[i], element + (i * ValueSize)); + } + } + + private static void WriteOwnedRecord( + IReadOnlyList> fields, + IntPtr destination) { - using (var scope = new AllocationScope()) + var count = fields.Count; + checked + { + _ = count * RecordEntrySize; + } + + Marshal.WriteByte(destination, (byte)ComponentValueKind.Record); + var vector = destination + ValuePayloadOffset; + ComponentValueNative.wasmtime_component_valrecord_new_uninit(vector, (nuint)count); + + var entry = Marshal.ReadIntPtr(vector + VectorDataOffset); + for (var i = 0; i < count; i++) + { + unsafe + { + new Span((void*)(entry + (i * RecordEntrySize)), RecordEntrySize).Clear(); + } + } + + for (var i = 0; i < count; i++) + { + var currentEntry = entry + (i * RecordEntrySize); + WriteOwnedName(fields[i].Key, currentEntry); + WriteOwned(fields[i].Value, currentEntry + RecordEntryValueOffset); + } + } + + private static IntPtr WriteOwnedBoxed(ComponentValue value) + { + unsafe + { + var boxed = stackalloc byte[ValueSize]; + var pointer = (IntPtr)boxed; + ClearValue(pointer); + + var initialized = false; + try + { + WriteOwned(value, pointer); + initialized = true; + + var owned = ComponentValueNative.wasmtime_component_val_new(pointer); + initialized = false; + return owned; + } + finally + { + if (initialized) + { + ComponentValueNative.wasmtime_component_val_delete(pointer); + } + } + } + } + + private static void ClearValue(IntPtr value) + { + unsafe { - var scratch = scope.Allocate(ValueSize); - Write(value, scratch, scope); - ComponentValueNative.wasmtime_component_val_clone(scratch, destination); + new Span((void*)value, ValueSize).Clear(); } } diff --git a/tests/ComponentValueLayoutTests.cs b/tests/ComponentValueLayoutTests.cs index 54a37439..5e63895a 100644 --- a/tests/ComponentValueLayoutTests.cs +++ b/tests/ComponentValueLayoutTests.cs @@ -284,12 +284,27 @@ public void OwnedWritesUseWasmtimesAllocator() ComponentValue.String("a host-allocated string"), ComponentValue.String(string.Empty), ComponentValue.Enum("warning"), + ComponentValue.List([]), ComponentValue.List(new[] { ComponentValue.String("a"), ComponentValue.String("b") }), + ComponentValue.Tuple(new[] { ComponentValue.S32(3), ComponentValue.String("tuple") }), ComponentValue.Record(new[] { new KeyValuePair("name", ComponentValue.String("x")), new KeyValuePair("value", ComponentValue.F64(1.5)), }), + ComponentValue.Record(new[] + { + new KeyValuePair("items", ComponentValue.List(new[] + { + ComponentValue.Some(ComponentValue.String("nested")), + ComponentValue.None(), + })), + new KeyValuePair("result", ComponentValue.Ok(ComponentValue.Tuple(new[] + { + ComponentValue.U16(12), + ComponentValue.String("done"), + }))), + }), ComponentValue.Some(ComponentValue.String("boxed")), ComponentValue.None(), ComponentValue.Ok(ComponentValue.S32(1)),