use std::sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
};
use std::time::Duration;
use serde_json::Value;
use wasmtime::component::{Component, Linker};
use wasmtime::{Config, Engine, ResourceLimiter, Store, Trap};
use crate::handler::HandlerLoadError;
use crate::handler_package::LoadedHandlerPack;
use crate::server::{InvokeError, InvokeErrorKind, InvokeHandler as InvokeHandlerTrait};
mod invoke_handler {
wasmtime::component::bindgen!({
path: "wit/rill-handler.wit",
world: "invoke-handler",
});
}
pub const CONFIGURE_FUEL: u64 = 10_000_000;
pub const INVOKE_FUEL: u64 = 100_000_000;
pub const MAX_MEMORY_BYTES: usize = 64 * 1024 * 1024;
pub const MAX_TABLE_ELEMENTS: u32 = 10_000;
pub const MAX_IO_BYTES: usize = 1024 * 1024;
pub const EPOCH_TICK_INTERVAL: Duration = Duration::from_secs(1);
pub const EPOCH_DEADLINE: u64 = 5;
#[cfg(test)]
static ACTIVE_EPOCH_TICKERS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
struct HostState;
impl ResourceLimiter for HostState {
fn memory_growing(
&mut self,
_current: usize,
desired: usize,
_max: Option<usize>,
) -> Result<bool, wasmtime::Error> {
Ok(desired <= MAX_MEMORY_BYTES)
}
fn table_growing(
&mut self,
_current: usize,
desired: usize,
_max: Option<usize>,
) -> Result<bool, wasmtime::Error> {
Ok(desired <= MAX_TABLE_ELEMENTS as usize)
}
}
struct WasmState {
store: Store<HostState>,
bindings: invoke_handler::InvokeHandler,
}
struct EpochTicker {
stop_flag: Arc<AtomicBool>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl EpochTicker {
fn start(engine: Engine) -> Self {
let stop_flag = Arc::new(AtomicBool::new(false));
let engine_for_thread = engine;
let stop_for_thread = Arc::clone(&stop_flag);
let handle = std::thread::spawn(move || {
#[cfg(test)]
ACTIVE_EPOCH_TICKERS.fetch_add(1, Ordering::SeqCst);
while !stop_for_thread.load(Ordering::Relaxed) {
std::thread::sleep(EPOCH_TICK_INTERVAL);
engine_for_thread.increment_epoch();
}
#[cfg(test)]
ACTIVE_EPOCH_TICKERS.fetch_sub(1, Ordering::SeqCst);
});
Self {
stop_flag,
handle: Some(handle),
}
}
}
impl Drop for EpochTicker {
fn drop(&mut self) {
self.stop_flag.store(true, Ordering::Relaxed);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
pub struct WasmInvokeHandler {
engine: Engine,
_ticker: EpochTicker,
state: Mutex<WasmState>,
}
impl WasmInvokeHandler {
pub fn new(pack: &LoadedHandlerPack, model_json: &Value) -> Result<Self, HandlerLoadError> {
let mut config = Config::new();
config.consume_fuel(true);
config.epoch_interruption(true);
config.max_wasm_stack(1024 * 1024);
let engine = Engine::new(&config)
.map_err(|e| HandlerLoadError::Init(format!("engine creation failed: {e}")))?;
let ticker = EpochTicker::start(engine.clone());
let component = Component::new(&engine, &pack.module)
.map_err(|e| HandlerLoadError::Init(format!("component compilation failed: {e}")))?;
let linker: Linker<HostState> = Linker::new(&engine);
let mut store = Store::new(&engine, HostState);
store.limiter(|state| state as &mut dyn ResourceLimiter);
store
.set_fuel(CONFIGURE_FUEL)
.map_err(|e| HandlerLoadError::Init(format!("failed to set instantiate fuel: {e}")))?;
store.set_epoch_deadline(EPOCH_DEADLINE);
let bindings = invoke_handler::InvokeHandler::instantiate(&mut store, &component, &linker)
.map_err(|e| HandlerLoadError::Init(format!("instantiation failed: {e}")))?;
store
.set_fuel(CONFIGURE_FUEL)
.map_err(|e| HandlerLoadError::Init(format!("failed to set metadata fuel: {e}")))?;
store.set_epoch_deadline(EPOCH_DEADLINE);
let metadata = bindings
.call_metadata(&mut store)
.map_err(|e| HandlerLoadError::Init(format!("metadata trap: {e}")))?;
if metadata.id != pack.manifest.id {
return Err(HandlerLoadError::MetadataMismatch(format!(
"guest id '{}' != manifest id '{}'",
metadata.id, pack.manifest.id
)));
}
if metadata.version != pack.manifest.version {
return Err(HandlerLoadError::MetadataMismatch(format!(
"guest version '{}' != manifest version '{}'",
metadata.version, pack.manifest.version
)));
}
if metadata.api_version != pack.manifest.handler_api_version {
return Err(HandlerLoadError::MetadataMismatch(format!(
"guest api version {} != manifest api version {}",
metadata.api_version, pack.manifest.handler_api_version
)));
}
let mut manifest_caps = pack.manifest.capabilities.clone();
manifest_caps.sort();
let mut metadata_caps = metadata.capabilities.clone();
metadata_caps.sort();
if manifest_caps != metadata_caps {
return Err(HandlerLoadError::MetadataMismatch(
"guest capabilities != manifest capabilities".into(),
));
}
let model_bytes = serde_json::to_vec(model_json)
.map_err(|e| HandlerLoadError::Init(format!("model serialization failed: {e}")))?;
if model_bytes.len() > MAX_IO_BYTES {
return Err(HandlerLoadError::Init("model JSON exceeds limit".into()));
}
store
.set_fuel(CONFIGURE_FUEL)
.map_err(|e| HandlerLoadError::Init(format!("failed to set configure fuel: {e}")))?;
store.set_epoch_deadline(EPOCH_DEADLINE);
let configure_result = bindings
.call_configure(&mut store, &model_bytes)
.map_err(|e| HandlerLoadError::Init(format!("configure trap: {e}")))?;
if let Err(handler_error) = configure_result {
let (variant, detail) = match handler_error {
invoke_handler::HandlerError::InvalidModel(s) => ("invalid-model", s),
invoke_handler::HandlerError::InvalidInput(s) => ("invalid-input", s),
invoke_handler::HandlerError::UnsupportedCapability(s) => {
("unsupported-capability", s)
}
invoke_handler::HandlerError::ExecutionFailed(s) => ("execution-failed", s),
};
return Err(HandlerLoadError::Init(format!(
"configure rejected model ({variant}): {detail}"
)));
}
Ok(Self {
engine,
_ticker: ticker,
state: Mutex::new(WasmState { store, bindings }),
})
}
#[allow(dead_code)]
pub fn engine(&self) -> &Engine {
&self.engine
}
}
impl std::fmt::Debug for WasmInvokeHandler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WasmInvokeHandler")
.field(
"epoch_ticker_running",
&!self._ticker.stop_flag.load(Ordering::Relaxed),
)
.finish_non_exhaustive()
}
}
impl InvokeHandlerTrait for WasmInvokeHandler {
fn invoke(&self, capability: &str, input: &Value) -> Result<Value, InvokeError> {
let input_bytes = serde_json::to_vec(input).map_err(|e| {
InvokeError::with_detail(
InvokeErrorKind::Internal,
format!("input serialization failed: {e}"),
)
})?;
if input_bytes.len() > MAX_IO_BYTES {
return Err(InvokeError::new(InvokeErrorKind::Internal));
}
let mut state = self.state.lock().map_err(|_| {
InvokeError::with_detail(InvokeErrorKind::Internal, "handler state mutex poisoned")
})?;
state.store.set_fuel(INVOKE_FUEL).map_err(|e| {
InvokeError::with_detail(
InvokeErrorKind::Internal,
format!("failed to set invoke fuel: {e}"),
)
})?;
state.store.set_epoch_deadline(EPOCH_DEADLINE);
let WasmState { store, bindings } = &mut *state;
let result = bindings
.call_invoke(store, capability, &input_bytes)
.map_err(|e| {
if let Some(trap) = e.downcast_ref::<Trap>()
&& matches!(trap, Trap::OutOfFuel | Trap::Interrupt)
{
return InvokeError::new(InvokeErrorKind::Timeout);
}
InvokeError::with_detail(InvokeErrorKind::Trap, format!("{e}"))
})?;
match result {
Ok(output_bytes) => {
if output_bytes.len() > MAX_IO_BYTES {
return Err(InvokeError::new(InvokeErrorKind::OutputTooLarge));
}
serde_json::from_slice(&output_bytes).map_err(|e| {
InvokeError::with_detail(
InvokeErrorKind::InvalidOutput,
format!("host-side JSON deserialisation failed: {e}"),
)
})
}
Err(handler_error) => {
let (kind, detail) = match handler_error {
invoke_handler::HandlerError::InvalidModel(s) => {
(InvokeErrorKind::InvalidModel, s)
}
invoke_handler::HandlerError::InvalidInput(s) => {
(InvokeErrorKind::InvalidInput, s)
}
invoke_handler::HandlerError::UnsupportedCapability(s) => {
(InvokeErrorKind::UnsupportedCapability, s)
}
invoke_handler::HandlerError::ExecutionFailed(s) => {
(InvokeErrorKind::ExecutionFailed, s)
}
};
Err(InvokeError::with_detail(kind, detail))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn memory_limiter_accepts_growth_within_max() {
let mut state = HostState;
assert!(
state
.memory_growing(0, MAX_MEMORY_BYTES, None)
.expect("memory_growing must not error")
);
assert!(
state
.memory_growing(0, MAX_MEMORY_BYTES - 1, None)
.expect("memory_growing must not error")
);
}
#[test]
fn memory_limiter_rejects_growth_exceeding_max() {
let mut state = HostState;
assert!(
!state
.memory_growing(MAX_MEMORY_BYTES - 1, MAX_MEMORY_BYTES + 1, None)
.expect("memory_growing must not error")
);
assert!(
!state
.memory_growing(0, MAX_MEMORY_BYTES * 2, None)
.expect("memory_growing must not error")
);
}
#[test]
fn table_limiter_accepts_growth_within_max() {
let mut state = HostState;
assert!(
state
.table_growing(0, MAX_TABLE_ELEMENTS as usize, None)
.expect("table_growing must not error")
);
assert!(
state
.table_growing(0, (MAX_TABLE_ELEMENTS - 1) as usize, None)
.expect("table_growing must not error")
);
}
#[test]
fn table_limiter_rejects_growth_exceeding_max() {
let mut state = HostState;
assert!(
!state
.table_growing(
(MAX_TABLE_ELEMENTS - 1) as usize,
(MAX_TABLE_ELEMENTS + 1) as usize,
None
)
.expect("table_growing must not error")
);
assert!(
!state
.table_growing(0, (MAX_TABLE_ELEMENTS * 2) as usize, None)
.expect("table_growing must not error")
);
}
use std::collections::BTreeMap;
use std::fs;
use std::path::PathBuf;
use std::sync::Arc;
use ed25519_dalek::{SigningKey, VerifyingKey};
use rill_runtime_protocol::{
HANDLER_API_VERSION, HANDLER_PACKAGE_FORMAT_VERSION, HandlerPackManifest,
};
use sha2::{Digest, Sha256};
static LIFECYCLE_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn active_epoch_ticker_count() -> usize {
ACTIVE_EPOCH_TICKERS.load(Ordering::SeqCst)
}
fn lifecycle_guard() -> std::sync::MutexGuard<'static, ()> {
LIFECYCLE_TEST_LOCK
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn wait_for_active_ticker_count(target: usize, timeout: std::time::Duration) -> bool {
let start = std::time::Instant::now();
loop {
if active_epoch_ticker_count() == target {
return true;
}
if start.elapsed() >= timeout {
return false;
}
std::thread::sleep(std::time::Duration::from_millis(10));
}
}
fn fixture_path(env_name: &str, fallback_relative: &str) -> Option<PathBuf> {
if let Ok(value) = std::env::var(env_name) {
let path = PathBuf::from(value);
assert!(
path.is_file(),
"{env_name} points to missing fixture: {}",
path.display()
);
return Some(path);
}
let fallback = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join(fallback_relative);
if fallback.is_file() {
return Some(fallback);
}
if std::env::var_os("RILL_RUN_WASM_FIXTURE_TESTS").is_some() {
panic!(
"{env_name} is not set and fallback fixture does not exist: {} \
(RILL_RUN_WASM_FIXTURE_TESTS=1 is set, so missing fixtures must fail)",
fallback.display()
);
}
None
}
fn echo_handler_component() -> Option<PathBuf> {
fixture_path("ECHO_HANDLER_WASM", "../../target/echo-handler.wasm")
}
fn metadata_loop_handler_component() -> Option<PathBuf> {
fixture_path(
"METADATA_LOOP_HANDLER_WASM",
"../../target/test-metadata-loop-handler.wasm",
)
}
fn build_echo_pack(module: &[u8], signing: &SigningKey) -> Vec<u8> {
let manifest = HandlerPackManifest {
format_version: HANDLER_PACKAGE_FORMAT_VERSION,
id: "rillml.echo.handler".into(),
version: env!("CARGO_PKG_VERSION").into(),
handler_api_version: HANDLER_API_VERSION,
min_runtime_version: env!("CARGO_PKG_VERSION").into(),
publisher_key_id: "wasm-test-key".into(),
capabilities: vec!["rillml.linearRegression.predict".into()],
module_sha256: hex::encode(Sha256::digest(module)),
module_size: module.len() as u64,
};
crate::build_signed_handler_pack(&manifest, module, signing).unwrap()
}
fn build_metadata_loop_pack(module: &[u8], signing: &SigningKey) -> Vec<u8> {
let manifest = HandlerPackManifest {
format_version: HANDLER_PACKAGE_FORMAT_VERSION,
id: "rillml.test.metadata-loop".into(),
version: env!("CARGO_PKG_VERSION").into(),
handler_api_version: HANDLER_API_VERSION,
min_runtime_version: env!("CARGO_PKG_VERSION").into(),
publisher_key_id: "wasm-test-key".into(),
capabilities: vec!["rillml.linearRegression.predict".into()],
module_sha256: hex::encode(Sha256::digest(module)),
module_size: module.len() as u64,
};
crate::build_signed_handler_pack(&manifest, module, signing).unwrap()
}
fn load_pack(pack_bytes: &[u8], verifying: &VerifyingKey) -> crate::LoadedHandlerPack {
let trust = crate::TrustStore(BTreeMap::from([("wasm-test-key".into(), *verifying)]));
let (loaded, _) =
crate::load_handler_pack(std::io::Cursor::new(pack_bytes), &trust).unwrap();
loaded
}
#[test]
fn normal_handler_drop_restores_active_ticker_count() {
let _guard = lifecycle_guard();
let component = match echo_handler_component() {
Some(path) => fs::read(&path).unwrap(),
None => {
eprintln!(
"skipping: echo handler component not built \
(set ECHO_HANDLER_WASM or RILL_RUN_WASM_FIXTURE_TESTS=1)"
);
return;
}
};
let signing = SigningKey::from_bytes(&[7; 32]);
let pack_bytes = build_echo_pack(&component, &signing);
let loaded = load_pack(&pack_bytes, &signing.verifying_key());
let baseline = active_epoch_ticker_count();
let model =
serde_json::json!({"kind": "linearRegression", "weights": [0.5], "intercept": 0.0});
let handler =
WasmInvokeHandler::new(&loaded, &model).expect("echo handler must load successfully");
assert!(
wait_for_active_ticker_count(baseline + 1, std::time::Duration::from_secs(3)),
"active ticker count did not reach {} after handler construction (got {}, baseline {})",
baseline + 1,
active_epoch_ticker_count(),
baseline
);
drop(handler);
assert!(
wait_for_active_ticker_count(baseline, std::time::Duration::from_secs(3)),
"active ticker count did not return to baseline {} after handler drop (got {})",
baseline,
active_epoch_ticker_count()
);
}
#[test]
fn metadata_loop_failure_restores_active_ticker_count() {
let _guard = lifecycle_guard();
let component = match metadata_loop_handler_component() {
Some(path) => fs::read(&path).unwrap(),
None => {
eprintln!(
"skipping: metadata-loop handler component not built \
(set METADATA_LOOP_HANDLER_WASM or RILL_RUN_WASM_FIXTURE_TESTS=1)"
);
return;
}
};
let signing = SigningKey::from_bytes(&[9; 32]);
let pack_bytes = build_metadata_loop_pack(&component, &signing);
let loaded = Arc::new(load_pack(&pack_bytes, &signing.verifying_key()));
let baseline = active_epoch_ticker_count();
let (tx, rx) = std::sync::mpsc::channel();
let worker_loaded = Arc::clone(&loaded);
let worker = std::thread::spawn(move || {
let result = WasmInvokeHandler::new(&worker_loaded, &serde_json::json!({}));
let _ = tx.send(result);
});
assert!(
wait_for_active_ticker_count(baseline + 1, std::time::Duration::from_secs(10)),
"metadata-loop constructor never started an epoch ticker \
(count stayed at {}, expected {} during construction)",
active_epoch_ticker_count(),
baseline + 1
);
let result = match rx.recv_timeout(std::time::Duration::from_secs(15)) {
Ok(result) => result,
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {
panic!(
"metadata-loop constructor did not terminate within test timeout (15s) \
— epoch interruption or worker exit logic may have regressed"
);
}
Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => {
let payload = worker.join().expect_err(
"metadata-loop constructor: rx disconnected but worker did not panic \
— this should be unreachable",
);
std::panic::resume_unwind(payload);
}
};
assert!(result.is_err(), "metadata-loop handler must fail to load");
worker
.join()
.expect("metadata-loop constructor thread panicked after sending result");
assert!(
wait_for_active_ticker_count(baseline, std::time::Duration::from_secs(3)),
"active ticker count did not return to baseline {} after metadata-loop failure (got {})",
baseline,
active_epoch_ticker_count()
);
println!("metadata-loop constructor timeout test: PASS");
drop(loaded);
}
#[test]
fn ticker_probe_is_available_to_internal_tests() {
let _guard = lifecycle_guard();
let baseline = active_epoch_ticker_count();
let _ = ACTIVE_EPOCH_TICKERS.load(Ordering::SeqCst);
assert_eq!(active_epoch_ticker_count(), baseline);
}
}