Skip to main content

numcodecs_wasm_host_reproducible/transform/
mod.rs

1use std::sync::OnceLock;
2
3use anyhow::{Context, Error, anyhow};
4use instcnt::PerfWitInterfaces;
5use numcodecs_wasm_host::NumcodecsWitInterfaces;
6use vecmap::VecMap;
7
8use crate::{logging::WasiLoggingInterface, stdio::WasiSandboxedStdioInterface};
9
10pub mod instcnt;
11pub mod nan;
12
13pub fn transform_wasm_component(wasm_component: impl Into<Vec<u8>>) -> Result<Vec<u8>, Error> {
14    let NumcodecsWitInterfaces {
15        package,
16        codec: codec_interface,
17        registry: registry_interface,
18        types: types_interface,
19        ..
20    } = NumcodecsWitInterfaces::get();
21
22    // create a new WAC composition graph with the WASI component packages
23    //  pre-registered and the numcodecs:wasm/perf interface pre-exported
24    let PreparedCompositionGraph {
25        graph: wac,
26        wasi: wasi_component_packages,
27    } = get_prepared_composition_graph()?;
28    let mut wac = wac.clone();
29
30    // parse and instantiate the root package, which exports numcodecs:abc/codec
31    let numcodecs_package = wac_graph::types::Package::from_bytes(
32        &format!("{}", package.name()),
33        package.version(),
34        wasm_component,
35        wac.types_mut(),
36    )?;
37
38    let numcodecs_world = &wac.types()[numcodecs_package.ty()];
39    let numcodecs_imports = extract_component_ports(&numcodecs_world.imports)?;
40
41    let numcodecs_package = wac.register_package(numcodecs_package)?;
42    let numcodecs_instance = wac.instantiate(numcodecs_package);
43
44    // list the imports that the linker will provide
45    let linker_provided_imports = [
46        &WasiSandboxedStdioInterface::get().stdio,
47        &WasiLoggingInterface::get().logging,
48        registry_interface,
49    ];
50    // trivial imports that require no linking
51    let trivial_imports = [types_interface];
52
53    // initialise the unresolved imports to the imports of the root package
54    let mut unresolved_imports = VecMap::new();
55    for import in &numcodecs_imports {
56        unresolved_imports
57            .entry(import.clone())
58            .or_insert_with(Vec::new)
59            .push(numcodecs_instance);
60    }
61
62    // track all non-root instances, which may fulfil imports
63    let mut package_instances = VecMap::new();
64
65    // iterate while not all unresolved imports have been satisfied
66    while let Some((unresolved_import, dependents)) = unresolved_imports.pop() {
67        // some imports are trivial and require no linking
68        if trivial_imports.contains(&&unresolved_import) {
69            continue;
70        }
71
72        // some imports will be provided by the linker
73        if linker_provided_imports.contains(&&unresolved_import) {
74            continue;
75        }
76
77        // resolve the import
78        // - either it has been instantiated already
79        // - or we instantiate the WASI component that will provide it
80        // - or it is unresolvable and we raise an error
81        let dependency_instance =
82            if let Some(dependency_instance) = package_instances.get(unresolved_import.package()) {
83                *dependency_instance
84            } else if let Some(component_package) = wasi_component_packages
85                .iter()
86                .find(|component_package| component_package.exports.contains(&unresolved_import))
87            {
88                let PackageWithPorts {
89                    package: component_package,
90                    imports: component_imports,
91                    exports: component_exports,
92                } = component_package;
93
94                // instantiate the component package
95                let component_instance = wac.instantiate(*component_package);
96
97                // require all imports of the component package to be resolved
98                for import in component_imports {
99                    unresolved_imports
100                        .entry(import.clone())
101                        .or_insert_with(Vec::new)
102                        .push(component_instance);
103                }
104
105                for export in component_exports {
106                    // register this instance's package so that its exports can later
107                    //  fulfil more imports
108                    package_instances.insert(export.package().clone(), component_instance);
109                }
110
111                component_instance
112            } else {
113                return Err(anyhow!(
114                    "WASM component requires unresolved import {unresolved_import}"
115                ));
116            };
117
118        let import_str = &format!("{unresolved_import}");
119        let dependency_export = wac.alias_instance_export(dependency_instance, import_str)?;
120        for dependent in dependents {
121            wac.set_instantiation_argument(dependent, import_str, dependency_export)?;
122        }
123    }
124
125    assert!(unresolved_imports.into_vec().is_empty());
126
127    // export the numcodecs:abc/codec interface
128    let numcodecs_codecs_str = &format!("{codec_interface}");
129    let numcodecs_codecs_export =
130        wac.alias_instance_export(numcodecs_instance, numcodecs_codecs_str)?;
131    wac.export(numcodecs_codecs_export, numcodecs_codecs_str)?;
132
133    // encode the WAC composition graph into a WASM component and validate it
134    let wasm = wac.encode(wac_graph::EncodeOptions {
135        define_components: true,
136        // we do our own validation right below
137        validate: false,
138        processor: None,
139    })?;
140
141    wasmparser::Validator::new_with_features(
142        wasmparser::WasmFeaturesInflated {
143            // MUST: float operations are required
144            //       (and our engine's transformations makes them deterministic)
145            floats: true,
146            // MUST: codecs and reproducible WASI are implemented as components
147            component_model: true,
148            ..crate::engine::DETERMINISTIC_WASM_MODULE_FEATURES
149        }
150        .into(),
151    )
152    .validate_all(&wasm)?;
153
154    Ok(wasm)
155}
156
157struct PreparedCompositionGraph {
158    graph: wac_graph::CompositionGraph,
159    wasi: Box<[PackageWithPorts]>,
160}
161
162fn get_prepared_composition_graph() -> Result<&'static PreparedCompositionGraph, Error> {
163    static PREPARED_COMPOSITION_GRAPH: OnceLock<Result<PreparedCompositionGraph, Error>> =
164        OnceLock::new();
165
166    let prepared_composition_graph = PREPARED_COMPOSITION_GRAPH.get_or_init(|| {
167        let PerfWitInterfaces {
168            perf: perf_interface,
169            ..
170        } = PerfWitInterfaces::get();
171
172        // create a new WAC composition graph
173        let mut wac = wac_graph::CompositionGraph::new();
174
175        // parse and register the WASI component packages
176        let wasi_component_packages =
177            register_wasi_component_packages(&mut wac)?.into_boxed_slice();
178
179        // create, register, and instantiate the numcodecs:wasm package
180        let numcodecs_wasm_perf_instance = instantiate_numcodecs_wasm_perf_package(&mut wac)?;
181
182        // export the numcodecs:wasm/perf interface
183        let numcodecs_wasm_perf_str = &format!("{perf_interface}");
184        let numcodecs_wasm_perf_export =
185            wac.alias_instance_export(numcodecs_wasm_perf_instance, numcodecs_wasm_perf_str)?;
186        wac.export(numcodecs_wasm_perf_export, numcodecs_wasm_perf_str)?;
187
188        Ok(PreparedCompositionGraph {
189            graph: wac,
190            wasi: wasi_component_packages,
191        })
192    });
193
194    match prepared_composition_graph {
195        Ok(prepared_composition_graph) => Ok(prepared_composition_graph),
196        Err(err) => Err(anyhow!(err)),
197    }
198}
199
200struct PackageWithPorts {
201    package: wac_graph::PackageId,
202    imports: Box<[wasm_component_layer::InterfaceIdentifier]>,
203    exports: Box<[wasm_component_layer::InterfaceIdentifier]>,
204}
205
206fn register_wasi_component_packages(
207    wac: &mut wac_graph::CompositionGraph,
208) -> Result<Vec<PackageWithPorts>, Error> {
209    const WASI_COMPONENTS: &[(&str, &[u8])] = &[(
210        "wasi-sandboxed:merged",
211        wasi_sandboxed_component_provider::MERGED_COMPONENT,
212    )];
213
214    let wasi_component_packages = WASI_COMPONENTS
215        .iter()
216        .map(|(component_name, component_bytes)| -> Result<_, Error> {
217            let component_package = wac_graph::types::Package::from_bytes(
218                component_name,
219                None,
220                Vec::from(*component_bytes),
221                wac.types_mut(),
222            )?;
223
224            let component_world = &wac.types()[component_package.ty()];
225
226            let component_imports = extract_component_ports(&component_world.imports)?;
227            let component_exports = extract_component_ports(&component_world.exports)?;
228
229            let component_package = wac.register_package(component_package)?;
230
231            Ok(PackageWithPorts {
232                package: component_package,
233                imports: component_imports.into_boxed_slice(),
234                exports: component_exports.into_boxed_slice(),
235            })
236        })
237        .collect::<Result<Vec<_>, _>>()?;
238
239    Ok(wasi_component_packages)
240}
241
242fn extract_component_ports(
243    ports: &indexmap::IndexMap<String, wac_graph::types::ItemKind>,
244) -> Result<Vec<wasm_component_layer::InterfaceIdentifier>, anyhow::Error> {
245    ports
246        .iter()
247        .filter_map(|(import, kind)| match kind {
248            wac_graph::types::ItemKind::Instance(_) => Some(
249                wasm_component_layer::InterfaceIdentifier::try_from(import.as_str()),
250            ),
251            _ => None,
252        })
253        .collect::<Result<Vec<_>, _>>()
254}
255
256fn instantiate_numcodecs_wasm_perf_package(
257    wac: &mut wac_graph::CompositionGraph,
258) -> Result<wac_graph::NodeId, Error> {
259    let PerfWitInterfaces {
260        perf: perf_interface,
261        ..
262    } = PerfWitInterfaces::get();
263
264    // create, register, and instantiate the numcodecs:wasm/perf package
265    let numcodecs_wasm_perf_package = wac_graph::types::Package::from_bytes(
266        &format!("{}", perf_interface.package().name()),
267        perf_interface.package().version(),
268        create_numcodecs_wasm_perf_component()?,
269        wac.types_mut(),
270    )?;
271
272    let numcodecs_wasm_perf_package = wac.register_package(numcodecs_wasm_perf_package)?;
273    let numcodecs_wasm_perf_instance = wac.instantiate(numcodecs_wasm_perf_package);
274
275    Ok(numcodecs_wasm_perf_instance)
276}
277
278fn create_numcodecs_wasm_perf_component() -> Result<Vec<u8>, Error> {
279    const ROOT: &str = "root";
280
281    let PerfWitInterfaces {
282        perf: perf_interface,
283        instruction_counter,
284    } = PerfWitInterfaces::get();
285
286    let mut module = create_numcodecs_wasm_perf_module();
287
288    let mut resolve = wit_parser::Resolve::new();
289
290    let interface = resolve.interfaces.alloc(wit_parser::Interface {
291        name: Some(String::from(perf_interface.name())),
292        types: indexmap::IndexMap::new(),
293        #[expect(clippy::iter_on_single_items)]
294        functions: [(
295            String::from(instruction_counter),
296            wit_parser::Function {
297                name: String::from(instruction_counter),
298                kind: wit_parser::FunctionKind::Freestanding,
299                params: Vec::new(),
300                result: Some(wit_parser::Type::U64),
301                docs: wit_parser::Docs { contents: None },
302                stability: wit_parser::Stability::Unknown,
303            },
304        )]
305        .into_iter()
306        .collect(),
307        docs: wit_parser::Docs { contents: None },
308        package: None, // The package is linked up below
309        stability: wit_parser::Stability::Unknown,
310    });
311
312    let package_name = wit_parser::PackageName {
313        namespace: String::from(perf_interface.package().name().namespace()),
314        name: String::from(perf_interface.package().name().name()),
315        version: perf_interface.package().version().cloned(),
316    };
317    let package = resolve.packages.alloc(wit_parser::Package {
318        name: package_name.clone(),
319        docs: wit_parser::Docs { contents: None },
320        #[expect(clippy::iter_on_single_items)]
321        interfaces: [(String::from(perf_interface.name()), interface)]
322            .into_iter()
323            .collect(),
324        worlds: indexmap::IndexMap::new(),
325    });
326    resolve.package_names.insert(package_name, package);
327
328    if let Some(interface) = resolve.interfaces.get_mut(interface) {
329        interface.package = Some(package);
330    }
331
332    let world = resolve.worlds.alloc(wit_parser::World {
333        name: String::from(ROOT),
334        imports: indexmap::IndexMap::new(),
335        #[expect(clippy::iter_on_single_items)]
336        exports: [(
337            wit_parser::WorldKey::Interface(interface),
338            wit_parser::WorldItem::Interface {
339                id: interface,
340                stability: wit_parser::Stability::Unknown,
341            },
342        )]
343        .into_iter()
344        .collect(),
345        package: None, // The package is linked up below
346        docs: wit_parser::Docs { contents: None },
347        includes: Vec::new(),
348        include_names: Vec::new(),
349        stability: wit_parser::Stability::Unknown,
350    });
351
352    let root_name = wit_parser::PackageName {
353        namespace: String::from(ROOT),
354        name: String::from("component"),
355        version: perf_interface.package().version().cloned(),
356    };
357    let root = resolve.packages.alloc(wit_parser::Package {
358        name: root_name.clone(),
359        docs: wit_parser::Docs { contents: None },
360        interfaces: indexmap::IndexMap::new(),
361        #[expect(clippy::iter_on_single_items)]
362        worlds: [(String::from(ROOT), world)].into_iter().collect(),
363    });
364    resolve.package_names.insert(root_name, root);
365
366    if let Some(world) = resolve.worlds.get_mut(world) {
367        world.package = Some(root);
368    }
369
370    wit_component::embed_component_metadata(
371        &mut module,
372        &resolve,
373        world,
374        wit_component::StringEncoding::UTF8,
375    )?;
376
377    let mut encoder = wit_component::ComponentEncoder::default()
378        .module(&module)
379        .context("wit_component::ComponentEncoder::module failed")?;
380
381    let component = encoder
382        .encode()
383        .context("wit_component::ComponentEncoder::encode failed")?;
384
385    Ok(component)
386}
387
388fn create_numcodecs_wasm_perf_module() -> Vec<u8> {
389    let PerfWitInterfaces {
390        perf: perf_interface,
391        instruction_counter,
392    } = PerfWitInterfaces::get();
393
394    let mut module = wasm_encoder::Module::new();
395
396    // Encode the type section with
397    //  types[0] = () -> i64
398    let mut types = wasm_encoder::TypeSection::new();
399    let ty0 = types.len();
400    types.ty().function([], [wasm_encoder::ValType::I64]);
401    module.section(&types);
402
403    // Encode the function section with
404    //  functions[0] = fn() -> i64 [ types[0] ]
405    let mut functions = wasm_encoder::FunctionSection::new();
406    let fn0 = functions.len();
407    functions.function(ty0);
408    module.section(&functions);
409
410    // Encode the export section with
411    //  {perf_interface}#{instruction_counter} = functions[0]
412    let mut exports = wasm_encoder::ExportSection::new();
413    exports.export(
414        &format!("{perf_interface}#{instruction_counter}"),
415        wasm_encoder::ExportKind::Func,
416        fn0,
417    );
418    module.section(&exports);
419
420    // Encode the code section.
421    let mut codes = wasm_encoder::CodeSection::new();
422    let mut fn0 = wasm_encoder::Function::new([]);
423    fn0.instruction(&wasm_encoder::Instruction::Unreachable);
424    fn0.instruction(&wasm_encoder::Instruction::End);
425    codes.function(&fn0);
426    module.section(&codes);
427
428    // Extract the encoded WASM bytes for this module
429    module.finish()
430}