diff --git a/src/engine.rs b/src/engine.rs index 9185552a..6f120a93 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -321,6 +321,49 @@ mod tests { Ok(()) } + #[test] + fn test_repeated_runs_match_fresh_runs() -> Result<()> { + for path in [ + "tests/fixtures/build/js_function.wasm", + "tests/fixtures/build/js_function_javy_plugin_v3.wasm", + ] + .map(Path::new) + { + let run_with = |engine: &Engine| -> Result { + let module = Module::from_file(engine, path)?; + let codec = if uses_msgpack_provider(&module) { + Codec::Messagepack + } else { + Codec::Json + }; + let input = BytesContainer::new( + BytesContainerType::Input, + codec, + include_bytes!("../tests/fixtures/input/js_function_input.json").to_vec(), + )?; + let result = run(FunctionRunParams { + function_path: path.to_path_buf(), + input, + export: DEFAULT_EXPORT, + module, + engine: engine.clone(), + scale_factor: 1.0, + profile_opts: None, + })?; + assert!(result.success); + Ok(serde_json::to_value(result)?) + }; + + let engine = new_engine()?; + let first = run_with(&engine)?; + assert_eq!(first, run_with(&engine)?); + assert_eq!(first, run_with(&new_engine()?)?); + assert_eq!(first, run_with(&engine)?); + } + + Ok(()) + } + #[test] fn test_exit_code_zero() -> Result<()> { let engine = new_engine()?; diff --git a/src/io.rs b/src/io.rs index 596cb9cb..ce2ea671 100644 --- a/src/io.rs +++ b/src/io.rs @@ -164,7 +164,8 @@ fn instantiate_imports( let mut mem_io_instance = None; if let Some(std_import) = module.std_import() { - let imported_module = Module::from_binary(engine, &std_import.bytes) + let imported_module = std_import + .module(engine) .unwrap_or_else(|_| panic!("Failed to load module {}", std_import.name)); let imported_module_instance = linker diff --git a/src/validated_module.rs b/src/validated_module.rs index 2d4b40e9..5224e1a6 100644 --- a/src/validated_module.rs +++ b/src/validated_module.rs @@ -1,13 +1,74 @@ -use std::borrow::Cow; +use std::{ + borrow::Cow, + collections::HashMap, + sync::{Mutex, PoisonError}, +}; use anyhow::{bail, Result}; use rust_embed::RustEmbed; -use wasmtime::Module; +use wasmtime::{Engine, Module}; #[derive(RustEmbed)] #[folder = "providers/"] struct StandardProviders; +static COMPILED_PROVIDERS: Mutex = Mutex::new(ProviderCache::new()); + +struct CompiledProviders { + engine: Engine, + modules: HashMap, +} + +struct ProviderCache(Option); + +impl ProviderCache { + const fn new() -> Self { + Self(None) + } + + fn get(&mut self, engine: &Engine, name: &str) -> Option { + match &mut self.0 { + Some(compiled) if Engine::same(&compiled.engine, engine) => { + compiled.modules.get(name).cloned() + } + cache => { + *cache = Some(CompiledProviders { + engine: engine.clone(), + modules: HashMap::new(), + }); + None + } + } + } + + fn insert(&mut self, engine: &Engine, name: &str, module: Module) -> Module { + match &mut self.0 { + Some(compiled) if Engine::same(&compiled.engine, engine) => compiled + .modules + .entry(name.to_string()) + .or_insert(module) + .clone(), + _ => module, + } + } +} + +fn cached_module( + cache: &Mutex, + engine: &Engine, + name: &str, + compile: impl FnOnce() -> Result, +) -> Result { + let lock = || cache.lock().unwrap_or_else(PoisonError::into_inner); + + if let Some(module) = lock().get(engine, name) { + return Ok(module); + } + + let module = compile()?; + Ok(lock().insert(engine, name, module)) +} + #[derive(Debug)] pub(crate) struct Provider { pub(crate) bytes: Cow<'static, [u8]>, @@ -15,6 +76,12 @@ pub(crate) struct Provider { } impl Provider { + pub(crate) fn module(&self, engine: &Engine) -> Result { + cached_module(&COMPILED_PROVIDERS, engine, &self.name, || { + Module::from_binary(engine, &self.bytes) + }) + } + pub(crate) fn is_mem_io_provider(&self) -> bool { let javy_plugin_version = self .name @@ -93,10 +160,133 @@ impl ValidatedModule { #[cfg(test)] mod tests { + use std::sync::Mutex; + use anyhow::Result; use wasmtime::{Engine, Module}; - use crate::validated_module::ValidatedModule; + use crate::validated_module::{ + cached_module, Provider, ProviderCache, StandardProviders, ValidatedModule, + }; + + fn provider(name: &str) -> Provider { + Provider { + bytes: StandardProviders::get(&format!("{name}.wasm")) + .unwrap() + .data, + name: name.into(), + } + } + + fn module( + cache: &Mutex, + engine: &Engine, + provider: &Provider, + ) -> Result { + cached_module(cache, engine, &provider.name, || { + Module::from_binary(engine, &provider.bytes) + }) + } + + #[test] + fn test_provider_cache_reuses_module_for_same_engine() -> Result<()> { + let engine = Engine::default(); + let provider = provider("shopify_function_v2"); + let cache = Mutex::new(ProviderCache::new()); + + let first = module(&cache, &engine, &provider)?; + let second = module(&cache, &engine.clone(), &provider)?; + + assert_eq!(first.image_range(), second.image_range()); + Ok(()) + } + + #[test] + fn test_provider_cache_keeps_each_provider() -> Result<()> { + let engine = Engine::default(); + let v1 = provider("shopify_function_v1"); + let v2 = provider("shopify_function_v2"); + let cache = Mutex::new(ProviderCache::new()); + + let first_v1 = module(&cache, &engine, &v1)?; + let first_v2 = module(&cache, &engine, &v2)?; + + assert_ne!(first_v1.image_range(), first_v2.image_range()); + assert_eq!( + first_v1.image_range(), + module(&cache, &engine, &v1)?.image_range() + ); + assert_eq!( + first_v2.image_range(), + module(&cache, &engine, &v2)?.image_range() + ); + Ok(()) + } + + #[test] + fn test_provider_cache_holds_only_most_recent_engine() -> Result<()> { + let first_engine = Engine::default(); + let second_engine = Engine::default(); + let provider = provider("shopify_function_v2"); + let cache = Mutex::new(ProviderCache::new()); + + let first = module(&cache, &first_engine, &provider)?; + let second = module(&cache, &second_engine, &provider)?; + assert!(Engine::same(second.engine(), &second_engine)); + assert_ne!(first.image_range(), second.image_range()); + + let first_again = module(&cache, &first_engine, &provider)?; + assert!(Engine::same(first_again.engine(), &first_engine)); + assert_ne!(first.image_range(), first_again.image_range()); + Ok(()) + } + + #[test] + fn test_provider_cache_serves_warm_lookups_during_compilation() -> Result<()> { + let engine = Engine::default(); + let v1 = provider("shopify_function_v1"); + let v2 = provider("shopify_function_v2"); + let cache = Mutex::new(ProviderCache::new()); + let warm_v1 = module(&cache, &engine, &v1)?; + + cached_module(&cache, &engine, &v2.name, || { + let during_compile = cache + .try_lock() + .expect("compilation holds the cache lock") + .get(&engine, &v1.name) + .expect("warm provider is cached"); + assert_eq!(warm_v1.image_range(), during_compile.image_range()); + Module::from_binary(&engine, &v2.bytes) + })?; + + Ok(()) + } + + #[test] + fn test_provider_cache_skips_insert_when_engine_changes_during_compilation() -> Result<()> { + let first_engine = Engine::default(); + let second_engine = Engine::default(); + let provider = provider("shopify_function_v2"); + let cache = Mutex::new(ProviderCache::new()); + let mut second = None; + + let first = cached_module(&cache, &first_engine, &provider.name, || { + second = Some(module(&cache, &second_engine, &provider)?); + Module::from_binary(&first_engine, &provider.bytes) + })?; + let second = second.unwrap(); + + assert!(Engine::same(first.engine(), &first_engine)); + assert_eq!( + second.image_range(), + module(&cache, &second_engine, &provider)?.image_range() + ); + assert_ne!( + first.image_range(), + module(&cache, &first_engine, &provider)?.image_range() + ); + Ok(()) + } #[test] fn test_module_with_just_wasi() -> Result<()> {