#![cfg(feature = "wasm")]
use std::sync::Arc;
use crate::mapreduce::wasm::{WasmModuleStore, WasmRawError};
pub const HOOK_ALLOC: &str = "hook_alloc";
pub const HOOK_PRECOMMIT: &str = "precommit";
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum PrecommitOutcome {
Accept(Vec<u8>),
Reject(String),
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum PrecommitError {
#[error("precommit hook module not found: {0}")]
ModuleNotFound(String),
#[error("precommit hook {module} failed: {message}")]
Runtime {
module: String,
message: String,
},
}
#[derive(Clone)]
pub struct PrecommitHooks {
store: Arc<WasmModuleStore>,
}
impl std::fmt::Debug for PrecommitHooks {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PrecommitHooks").finish_non_exhaustive()
}
}
impl PrecommitHooks {
#[must_use]
pub fn new(store: Arc<WasmModuleStore>) -> Self {
Self { store }
}
pub fn register(
&self,
id: impl Into<String>,
bytes: &[u8],
) -> Result<(), crate::datatypes::keyfun::KeyFunError> {
self.store.register(id.into(), bytes).map_err(|e| {
crate::datatypes::keyfun::KeyFunError::Runtime {
module: "precommit".into(),
message: e.to_string(),
}
})
}
#[must_use]
pub fn contains(&self, id: &str) -> bool {
self.store.contains(id)
}
pub fn run(&self, module_id: &str, value: &[u8]) -> Result<PrecommitOutcome, PrecommitError> {
if module_id.is_empty() || !self.store.contains(module_id) {
return Err(PrecommitError::ModuleNotFound(module_id.to_string()));
}
match self
.store
.run_module_raw(module_id, value, HOOK_ALLOC, HOOK_PRECOMMIT)
{
Ok(out) => Ok(PrecommitOutcome::Accept(out)),
Err(WasmRawError::Status { message, .. }) => Ok(PrecommitOutcome::Reject(message)),
Err(WasmRawError::NotFound) => {
Err(PrecommitError::ModuleNotFound(module_id.to_string()))
}
Err(e) => Err(PrecommitError::Runtime {
module: module_id.to_string(),
message: format!("{e:?}"),
}),
}
}
}
impl crate::router::PrecommitRunner for PrecommitHooks {
fn run(&self, module_id: &str, value: &[u8]) -> Result<Vec<u8>, crate::router::PrecommitVeto> {
match PrecommitHooks::run(self, module_id, value) {
Ok(PrecommitOutcome::Accept(v)) => Ok(v),
Ok(PrecommitOutcome::Reject(reason)) => {
Err(crate::router::PrecommitVeto::Rejected(reason))
}
Err(e) => Err(crate::router::PrecommitVeto::Error(e.to_string())),
}
}
}
pub const HOOK_ALLOC_POSTCOMMIT: &str = "hook_alloc";
pub const HOOK_POSTCOMMIT: &str = "postcommit";
#[derive(Clone)]
pub struct PostcommitHooks {
store: Arc<WasmModuleStore>,
}
impl std::fmt::Debug for PostcommitHooks {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PostcommitHooks").finish_non_exhaustive()
}
}
impl PostcommitHooks {
#[must_use]
pub fn new(store: Arc<WasmModuleStore>) -> Self {
Self { store }
}
pub fn register(
&self,
id: impl Into<String>,
bytes: &[u8],
) -> Result<(), crate::datatypes::keyfun::KeyFunError> {
self.store.register(id.into(), bytes).map_err(|e| {
crate::datatypes::keyfun::KeyFunError::Runtime {
module: "postcommit".into(),
message: e.to_string(),
}
})
}
#[must_use]
pub fn contains(&self, id: &str) -> bool {
self.store.contains(id)
}
pub fn run(&self, module_id: &str, value: &[u8]) {
if module_id.is_empty() {
return;
}
if !self.store.contains(module_id) {
tracing::warn!(module = %module_id, "postcommit hook module not found");
return;
}
match self
.store
.run_module_raw(module_id, value, HOOK_ALLOC_POSTCOMMIT, HOOK_POSTCOMMIT)
{
Ok(_) => {}
Err(WasmRawError::Status { code, message }) => {
tracing::warn!(
module = %module_id,
status = code,
%message,
"postcommit hook returned a non-zero status"
);
}
Err(e) => {
tracing::warn!(
module = %module_id,
error = ?e,
"postcommit hook failed to run"
);
}
}
}
}
impl crate::router::PostcommitRunner for PostcommitHooks {
fn run(&self, module_id: &str, value: &[u8]) {
PostcommitHooks::run(self, module_id, value);
}
}
#[cfg(test)]
mod tests {
use super::*;
const ACCEPT_WAT: &str = r#"
(module
(memory (export "memory") 1)
(global $heap_top (mut i32) (i32.const 1024))
(func $alloc_inner (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 "hook_alloc") (param $len i32) (result i32)
(call $alloc_inner (local.get $len)))
(func (export "precommit")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
(local $out_buf i32)
(local.set $out_buf (call $alloc_inner (local.get $in_len)))
(memory.copy (local.get $out_buf) (local.get $in_ptr) (local.get $in_len))
(i32.store (local.get $out_ptr_ptr) (local.get $out_buf))
(i32.store (local.get $out_len_ptr) (local.get $in_len))
(i32.const 0)))
"#;
const REJECT_WAT: &str = r#"
(module
(memory (export "memory") 1)
(data (i32.const 2048) "denied")
(global $heap_top (mut i32) (i32.const 1024))
(func (export "hook_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 "precommit")
(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 6))
(i32.const 1)))
"#;
fn hooks() -> PrecommitHooks {
PrecommitHooks::new(Arc::new(WasmModuleStore::new().expect("wasm store")))
}
#[test]
fn accept_hook_passes_the_value_through() {
let h = hooks();
h.register("ok", ACCEPT_WAT.as_bytes()).expect("register");
assert_eq!(
h.run("ok", b"payload").expect("run"),
PrecommitOutcome::Accept(b"payload".to_vec())
);
}
#[test]
fn reject_hook_vetoes_with_a_reason() {
let h = hooks();
h.register("no", REJECT_WAT.as_bytes()).expect("register");
assert_eq!(
h.run("no", b"payload").expect("run"),
PrecommitOutcome::Reject("denied".to_string())
);
}
#[test]
fn unknown_module_is_not_found() {
let h = hooks();
assert!(matches!(
h.run("absent", b"x"),
Err(PrecommitError::ModuleNotFound(_))
));
}
const POSTCOMMIT_OK_WAT: &str = r#"
(module
(memory (export "memory") 1)
(global $heap_top (mut i32) (i32.const 1024))
(func $alloc_inner (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 "hook_alloc") (param $len i32) (result i32)
(call $alloc_inner (local.get $len)))
(func (export "postcommit")
(param $in_ptr i32) (param $in_len i32)
(param $out_ptr_ptr i32) (param $out_len_ptr i32)
(result i32)
(local $out_buf i32)
(local.set $out_buf (call $alloc_inner (local.get $in_len)))
(memory.copy (local.get $out_buf) (local.get $in_ptr) (local.get $in_len))
(i32.store (local.get $out_ptr_ptr) (local.get $out_buf))
(i32.store (local.get $out_len_ptr) (local.get $in_len))
(i32.const 0)))
"#;
const POSTCOMMIT_FAIL_WAT: &str = r#"
(module
(memory (export "memory") 1)
(data (i32.const 2048) "side effect failed")
(global $heap_top (mut i32) (i32.const 1024))
(func (export "hook_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 "postcommit")
(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 19))
(i32.const 1)))
"#;
fn postcommit_hooks() -> PostcommitHooks {
PostcommitHooks::new(Arc::new(WasmModuleStore::new().expect("wasm store")))
}
#[test]
fn postcommit_hook_runs_over_the_committed_value() {
let h = postcommit_hooks();
h.register("notify", POSTCOMMIT_OK_WAT.as_bytes())
.expect("register");
h.run("notify", b"committed-value");
}
#[test]
fn postcommit_hook_non_zero_status_does_not_panic_or_propagate() {
let h = postcommit_hooks();
h.register("flaky", POSTCOMMIT_FAIL_WAT.as_bytes())
.expect("register");
h.run("flaky", b"committed-value");
}
#[test]
fn postcommit_missing_module_does_not_panic_or_propagate() {
let h = postcommit_hooks();
h.run("absent", b"committed-value");
}
#[test]
fn postcommit_runner_trait_delegates_to_run() {
use crate::router::PostcommitRunner;
let h = postcommit_hooks();
h.register("notify", POSTCOMMIT_OK_WAT.as_bytes())
.expect("register");
PostcommitRunner::run(&h, "notify", b"v");
}
}