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: "../rill-handler-api/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;
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 || {
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();
}
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();
}
}
}
#[doc(hidden)]
pub fn active_epoch_ticker_count() -> usize {
ACTIVE_EPOCH_TICKERS.load(Ordering::SeqCst)
}
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")
);
}
}