1use std::sync::{
4 Arc, Mutex,
5 atomic::{AtomicBool, Ordering},
6};
7use std::time::Duration;
8
9use rill_handler_api::v2::{MAX_EVENT_BYTES, MAX_OUTPUT_BYTES, MAX_STATE_BYTES};
10use serde_json::Value;
11use wasmtime::component::{Component, Linker};
12use wasmtime::{Config, Engine, ResourceLimiter, Store, Trap};
13
14use crate::{
15 StatefulHandlerErrorKindV2, StatefulHandlerErrorV2, StatefulHandlerMetadataV2,
16 StatefulHandlerResultV2, StatefulHandlerV2,
17};
18
19mod bindings {
20 wasmtime::component::bindgen!({
21 path: "wit-v2/rill-handler.wit",
22 world: "stateful-handler",
23 });
24}
25
26pub const STATEFUL_CONFIGURE_FUEL: u64 = 10_000_000;
27pub const STATEFUL_HANDLE_FUEL: u64 = 100_000_000;
28pub const STATEFUL_EPOCH_TICK_INTERVAL: Duration = Duration::from_secs(1);
29pub const STATEFUL_EPOCH_DEADLINE: u64 = 5;
30
31struct HostStateV2;
32
33impl ResourceLimiter for HostStateV2 {
34 fn memory_growing(
35 &mut self,
36 _current: usize,
37 desired: usize,
38 _maximum: Option<usize>,
39 ) -> Result<bool, wasmtime::Error> {
40 Ok(desired <= crate::MAX_MEMORY_BYTES)
41 }
42
43 fn table_growing(
44 &mut self,
45 _current: usize,
46 desired: usize,
47 _maximum: Option<usize>,
48 ) -> Result<bool, wasmtime::Error> {
49 Ok(desired <= crate::MAX_TABLE_ELEMENTS as usize)
50 }
51}
52
53struct WasmStateV2 {
54 store: Store<HostStateV2>,
55 bindings: bindings::StatefulHandler,
56}
57
58struct EpochTickerV2 {
59 stop: Arc<AtomicBool>,
60 handle: Option<std::thread::JoinHandle<()>>,
61}
62
63impl EpochTickerV2 {
64 fn start(engine: Engine) -> Self {
65 let stop = Arc::new(AtomicBool::new(false));
66 let thread_stop = Arc::clone(&stop);
67 let handle = std::thread::spawn(move || {
68 while !thread_stop.load(Ordering::Relaxed) {
69 std::thread::sleep(STATEFUL_EPOCH_TICK_INTERVAL);
70 engine.increment_epoch();
71 }
72 });
73 Self {
74 stop,
75 handle: Some(handle),
76 }
77 }
78}
79
80impl Drop for EpochTickerV2 {
81 fn drop(&mut self) {
82 self.stop.store(true, Ordering::Relaxed);
83 if let Some(handle) = self.handle.take() {
84 let _ = handle.join();
85 }
86 }
87}
88
89pub struct WasmStatefulHandlerV2 {
93 metadata: StatefulHandlerMetadataV2,
94 _engine: Engine,
95 _ticker: EpochTickerV2,
96 state: Mutex<WasmStateV2>,
97}
98
99impl WasmStatefulHandlerV2 {
100 pub fn new(
103 component_bytes: &[u8],
104 expected_metadata: StatefulHandlerMetadataV2,
105 model_json: &Value,
106 ) -> Result<Self, StatefulHandlerErrorV2> {
107 let mut config = Config::new();
108 config.consume_fuel(true);
109 config.epoch_interruption(true);
110 config.max_wasm_stack(1024 * 1024);
111 let engine = Engine::new(&config).map_err(|error| {
112 StatefulHandlerErrorV2::with_detail(
113 StatefulHandlerErrorKindV2::Internal,
114 error.to_string(),
115 )
116 })?;
117 let ticker = EpochTickerV2::start(engine.clone());
118 let component = Component::new(&engine, component_bytes).map_err(|error| {
119 StatefulHandlerErrorV2::with_detail(
120 StatefulHandlerErrorKindV2::InvalidModel,
121 error.to_string(),
122 )
123 })?;
124 let linker: Linker<HostStateV2> = Linker::new(&engine);
126 let mut store = Store::new(&engine, HostStateV2);
127 store.limiter(|state| state as &mut dyn ResourceLimiter);
128 set_budget(&mut store, STATEFUL_CONFIGURE_FUEL)?;
129 let bindings = bindings::StatefulHandler::instantiate(&mut store, &component, &linker)
130 .map_err(map_load_trap)?;
131
132 set_budget(&mut store, STATEFUL_CONFIGURE_FUEL)?;
133 let guest = bindings.call_metadata(&mut store).map_err(map_load_trap)?;
134 let actual = StatefulHandlerMetadataV2 {
135 id: guest.id,
136 version: guest.version,
137 api_version: guest.api_version,
138 capabilities: guest.capabilities,
139 state_schema_version: guest.state_schema_version,
140 };
141 if actual != expected_metadata {
142 return Err(StatefulHandlerErrorV2::new(
143 StatefulHandlerErrorKindV2::MetadataMismatch,
144 ));
145 }
146
147 let model_bytes = serde_json::to_vec(model_json).map_err(|error| {
148 StatefulHandlerErrorV2::with_detail(
149 StatefulHandlerErrorKindV2::InvalidModel,
150 error.to_string(),
151 )
152 })?;
153 if model_bytes.len() > MAX_EVENT_BYTES {
154 return Err(StatefulHandlerErrorV2::new(
155 StatefulHandlerErrorKindV2::InvalidModel,
156 ));
157 }
158 set_budget(&mut store, STATEFUL_CONFIGURE_FUEL)?;
159 let configured = bindings
160 .call_configure(&mut store, &model_bytes)
161 .map_err(map_load_trap)?;
162 if let Err(error) = configured {
163 return Err(map_guest_error(error));
164 }
165
166 Ok(Self {
167 metadata: actual,
168 _engine: engine,
169 _ticker: ticker,
170 state: Mutex::new(WasmStateV2 { store, bindings }),
171 })
172 }
173}
174
175impl std::fmt::Debug for WasmStatefulHandlerV2 {
176 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
177 f.debug_struct("WasmStatefulHandlerV2")
178 .field("metadata", &self.metadata)
179 .finish_non_exhaustive()
180 }
181}
182
183impl StatefulHandlerV2 for WasmStatefulHandlerV2 {
184 fn metadata(&self) -> &StatefulHandlerMetadataV2 {
185 &self.metadata
186 }
187
188 fn handle(
189 &self,
190 event_json: &[u8],
191 current_state: &[u8],
192 deterministic_seed: Option<u64>,
193 ) -> Result<StatefulHandlerResultV2, StatefulHandlerErrorV2> {
194 if event_json.len() > MAX_EVENT_BYTES || current_state.len() > MAX_STATE_BYTES {
195 return Err(StatefulHandlerErrorV2::new(
196 StatefulHandlerErrorKindV2::InvalidEvent,
197 ));
198 }
199 let mut state = self
200 .state
201 .lock()
202 .map_err(|_| StatefulHandlerErrorV2::new(StatefulHandlerErrorKindV2::Internal))?;
203 set_budget(&mut state.store, STATEFUL_HANDLE_FUEL)?;
204 let WasmStateV2 { store, bindings } = &mut *state;
205 let result = bindings
206 .call_handle(store, event_json, current_state, deterministic_seed)
207 .map_err(map_call_trap)?;
208 let result = result.map_err(map_guest_error)?;
209 if result.output_json.len() > MAX_OUTPUT_BYTES {
210 return Err(StatefulHandlerErrorV2::new(
211 StatefulHandlerErrorKindV2::OutputTooLarge,
212 ));
213 }
214 if result.next_state.len() > MAX_STATE_BYTES {
215 return Err(StatefulHandlerErrorV2::new(
216 StatefulHandlerErrorKindV2::InvalidState,
217 ));
218 }
219 let output = serde_json::from_slice(&result.output_json).map_err(|error| {
220 StatefulHandlerErrorV2::with_detail(
221 StatefulHandlerErrorKindV2::InvalidOutput,
222 error.to_string(),
223 )
224 })?;
225 Ok(StatefulHandlerResultV2 {
226 output,
227 next_state: result.next_state,
228 })
229 }
230}
231
232fn set_budget(store: &mut Store<HostStateV2>, fuel: u64) -> Result<(), StatefulHandlerErrorV2> {
233 store.set_fuel(fuel).map_err(|error| {
234 StatefulHandlerErrorV2::with_detail(StatefulHandlerErrorKindV2::Internal, error.to_string())
235 })?;
236 store.set_epoch_deadline(STATEFUL_EPOCH_DEADLINE);
237 Ok(())
238}
239
240fn map_load_trap(error: wasmtime::Error) -> StatefulHandlerErrorV2 {
241 if let Some(trap) = error.downcast_ref::<Trap>()
242 && matches!(trap, Trap::OutOfFuel | Trap::Interrupt)
243 {
244 return StatefulHandlerErrorV2::new(StatefulHandlerErrorKindV2::Timeout);
245 }
246 StatefulHandlerErrorV2::with_detail(StatefulHandlerErrorKindV2::Trap, error.to_string())
247}
248
249fn map_call_trap(error: wasmtime::Error) -> StatefulHandlerErrorV2 {
250 map_load_trap(error)
251}
252
253fn map_guest_error(error: bindings::HandlerErrorV2) -> StatefulHandlerErrorV2 {
254 let (kind, detail) = match error {
255 bindings::HandlerErrorV2::InvalidModel(detail) => {
256 (StatefulHandlerErrorKindV2::InvalidModel, detail)
257 }
258 bindings::HandlerErrorV2::InvalidEvent(detail) => {
259 (StatefulHandlerErrorKindV2::InvalidEvent, detail)
260 }
261 bindings::HandlerErrorV2::InvalidState(detail) => {
262 (StatefulHandlerErrorKindV2::InvalidState, detail)
263 }
264 bindings::HandlerErrorV2::IncompatibleVersion(detail) => {
265 (StatefulHandlerErrorKindV2::IncompatibleVersion, detail)
266 }
267 bindings::HandlerErrorV2::ExecutionFailed(detail) => {
268 (StatefulHandlerErrorKindV2::Internal, detail)
269 }
270 };
271 StatefulHandlerErrorV2::with_detail(kind, detail)
272}