use std::sync::Arc;
use crate::datatypes::keyfun::KeyFunError;
use crate::mapreduce::wasm::{WasmLimits, WasmModuleStore, WasmRawError, WasmStoreError};
pub const KEYFUN_ALLOC: &str = "keyfun_alloc";
pub const KEYFUN_ROUTE: &str = "keyfun_route";
#[derive(Clone)]
pub struct WasmKeyfunStore {
inner: Arc<WasmModuleStore>,
}
impl std::fmt::Debug for WasmKeyfunStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WasmKeyfunStore")
.field("modules", &self.inner.module_ids())
.finish()
}
}
impl WasmKeyfunStore {
pub fn new() -> Result<Self, WasmStoreError> {
Ok(Self {
inner: Arc::new(WasmModuleStore::new()?),
})
}
pub fn with_limits(limits: WasmLimits) -> Result<Self, WasmStoreError> {
Ok(Self {
inner: Arc::new(WasmModuleStore::with_limits(limits)?),
})
}
#[must_use]
pub fn from_module_store(inner: Arc<WasmModuleStore>) -> Self {
Self { inner }
}
#[must_use]
pub fn module_store(&self) -> &Arc<WasmModuleStore> {
&self.inner
}
pub fn register(&self, id: impl Into<String>, bytes: &[u8]) -> Result<(), WasmStoreError> {
self.inner.register(id, bytes)
}
#[must_use]
pub fn contains(&self, id: &str) -> bool {
self.inner.contains(id)
}
#[must_use]
pub fn count(&self) -> usize {
self.inner.count()
}
#[must_use]
pub fn module_ids(&self) -> Vec<String> {
self.inner.module_ids()
}
#[must_use]
pub fn frame_input(bucket: &[u8], key: &[u8]) -> Vec<u8> {
let mut buf = Vec::with_capacity(8 + bucket.len() + key.len());
let blen = u32::try_from(bucket.len()).unwrap_or(u32::MAX);
let klen = u32::try_from(key.len()).unwrap_or(u32::MAX);
buf.extend_from_slice(&blen.to_le_bytes());
buf.extend_from_slice(bucket);
buf.extend_from_slice(&klen.to_le_bytes());
buf.extend_from_slice(key);
buf
}
pub fn route_bytes(
&self,
module_id: &str,
bucket: &[u8],
key: &[u8],
) -> Result<Vec<u8>, KeyFunError> {
if module_id.is_empty() {
return Err(KeyFunError::ModuleNotFound(module_id.to_string()));
}
let input = Self::frame_input(bucket, key);
self.inner
.run_module_raw(module_id, &input, KEYFUN_ALLOC, KEYFUN_ROUTE)
.map_err(|e| map_raw_error(module_id, e))
}
}
fn map_raw_error(module_id: &str, e: WasmRawError) -> KeyFunError {
match e {
WasmRawError::NotFound => KeyFunError::ModuleNotFound(module_id.to_string()),
WasmRawError::MemoryLimit => KeyFunError::MemoryLimit(module_id.to_string()),
WasmRawError::Timeout => KeyFunError::Runtime {
module: module_id.to_string(),
message: "execution timed out (fuel or wall-clock deadline)".into(),
},
WasmRawError::Status { code, message } => KeyFunError::Runtime {
module: module_id.to_string(),
message: format!("module returned status {code}: {message}"),
},
WasmRawError::Runtime(m) => KeyFunError::Runtime {
module: module_id.to_string(),
message: m,
},
WasmRawError::InputTooLarge => KeyFunError::Runtime {
module: module_id.to_string(),
message: "framed input length exceeds i32".into(),
},
}
}
#[cfg(test)]
mod tests {
use super::*;
const REVERSE_KEY_WAT: &str = r#"
(module
(memory (export "memory") 1)
(global $heap_top (mut i32) (i32.const 1024))
(func $alloc (param $len i32) (result i32)
(local $ptr i32)
(local.set $ptr (global.get $heap_top))
(global.set $heap_top
(i32.add (global.get $heap_top) (local.get $len)))
(local.get $ptr))
(func (export "keyfun_alloc") (param $len i32) (result i32)
(call $alloc (local.get $len)))
(func (export "keyfun_route")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
(local $blen i32)
(local $klen i32)
(local $kstart i32)
(local $out i32)
(local $i i32)
;; blen = load u32 at in_ptr
(local.set $blen (i32.load (local.get $in_ptr)))
;; klen = load u32 at in_ptr + 4 + blen
(local.set $klen
(i32.load
(i32.add (local.get $in_ptr)
(i32.add (i32.const 4) (local.get $blen)))))
;; kstart = in_ptr + 4 + blen + 4
(local.set $kstart
(i32.add (local.get $in_ptr)
(i32.add (i32.const 8) (local.get $blen))))
;; out = alloc(klen)
(local.set $out (call $alloc (local.get $klen)))
;; reverse copy
(local.set $i (i32.const 0))
(block $done
(loop $loop
(br_if $done (i32.ge_s (local.get $i) (local.get $klen)))
(i32.store8
(i32.add (local.get $out) (local.get $i))
(i32.load8_u
(i32.add (local.get $kstart)
(i32.sub (i32.sub (local.get $klen) (local.get $i))
(i32.const 1)))))
(local.set $i (i32.add (local.get $i) (i32.const 1)))
(br $loop)))
(i32.store (local.get $out_ptr_ptr) (local.get $out))
(i32.store (local.get $out_len_ptr) (local.get $klen))
(i32.const 0)))
"#;
const TRAP_WAT: &str = r#"
(module
(memory (export "memory") 1)
(func (export "keyfun_alloc") (param $len i32) (result i32)
(i32.const 1024))
(func (export "keyfun_route")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
unreachable))
"#;
const MEMORY_HOG_WAT: &str = r#"
(module
(memory (export "memory") 1)
(func (export "keyfun_alloc") (param $len i32) (result i32)
(i32.const 1024))
(func (export "keyfun_route")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
(if (i32.eq (memory.grow (i32.const 4096)) (i32.const -1))
(then unreachable))
(i32.const 0)))
"#;
const STATUS_ERR_WAT: &str = r#"
(module
(memory (export "memory") 1)
(data (i32.const 2048) "boom")
(func (export "keyfun_alloc") (param $len i32) (result i32)
(i32.const 1024))
(func (export "keyfun_route")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
(i32.store (local.get $out_ptr_ptr) (i32.const 2048))
(i32.store (local.get $out_len_ptr) (i32.const 4))
(i32.const 7)))
"#;
#[test]
fn reverse_key_wat_routes() {
let store = WasmKeyfunStore::new().expect("store");
store
.register("reverse", REVERSE_KEY_WAT.as_bytes())
.expect("register");
let out = store
.route_bytes("reverse", b"users", b"alice")
.expect("ok");
assert_eq!(out, b"ecila");
let out2 = store
.route_bytes("reverse", b"orders", b"alice")
.expect("ok");
assert_eq!(out2, b"ecila");
}
#[test]
fn empty_module_id_is_not_found() {
let store = WasmKeyfunStore::new().expect("store");
let err = store.route_bytes("", b"b", b"k").expect_err("err");
assert!(matches!(err, KeyFunError::ModuleNotFound(_)));
}
#[test]
fn unregistered_module_is_not_found() {
let store = WasmKeyfunStore::new().expect("store");
let err = store.route_bytes("nope", b"b", b"k").expect_err("err");
assert!(matches!(err, KeyFunError::ModuleNotFound(ref s) if s == "nope"));
}
#[test]
fn trapping_module_is_runtime_error() {
let store = WasmKeyfunStore::new().expect("store");
store
.register("trap", TRAP_WAT.as_bytes())
.expect("register");
let err = store.route_bytes("trap", b"b", b"k").expect_err("err");
assert!(matches!(err, KeyFunError::Runtime { .. }));
}
#[test]
fn oversize_memory_is_memory_limit_error() {
let store = WasmKeyfunStore::new().expect("store");
store
.register("hog", MEMORY_HOG_WAT.as_bytes())
.expect("register");
let err = store.route_bytes("hog", b"b", b"k").expect_err("err");
assert!(matches!(err, KeyFunError::MemoryLimit(_)));
}
#[test]
fn nonzero_status_is_runtime_error_with_message() {
let store = WasmKeyfunStore::new().expect("store");
store
.register("status", STATUS_ERR_WAT.as_bytes())
.expect("register");
let err = store.route_bytes("status", b"b", b"k").expect_err("err");
match err {
KeyFunError::Runtime { ref message, .. } => {
assert!(message.contains("status 7"), "got {message}");
assert!(message.contains("boom"), "got {message}");
}
other => panic!("expected Runtime, got {other:?}"),
}
}
#[test]
fn frame_input_layout_is_length_prefixed() {
let framed = WasmKeyfunStore::frame_input(b"bk", b"key");
assert_eq!(
framed,
[
&2u32.to_le_bytes()[..],
b"bk",
&3u32.to_le_bytes()[..],
b"key"
]
.concat()
);
}
#[test]
fn shares_one_module_store_with_mapreduce() {
let shared = std::sync::Arc::new(WasmModuleStore::new().expect("module store"));
let keyfun = WasmKeyfunStore::from_module_store(shared.clone());
keyfun
.register("reverse", REVERSE_KEY_WAT.as_bytes())
.expect("register");
assert!(shared.contains("reverse"));
let out = keyfun.route_bytes("reverse", b"b", b"abc").expect("route");
assert_eq!(out, b"cba");
assert_eq!(keyfun.module_store().count(), 1);
}
}