Skip to main content

polyhook_core/
wasm.rs

1/// WASM ABI layer.
2///
3/// Memory protocol
4/// ---------------
5/// All strings cross the WASM boundary as length-prefixed blobs:
6///   [4 bytes LE i32 = payload length][payload bytes...]
7///
8/// Callers must:
9///   1. Call `alloc(len)` to get a pointer.
10///   2. Write their payload into WASM memory at that pointer.
11///   3. Call `parse` or `serialize` with (ptr, len).
12///   4. Read the 4-byte LE length from the returned pointer, then the payload.
13///   5. Call `dealloc` on both the input buffer and the output buffer.
14use std::cell::RefCell;
15
16use crate::parse::parse_event;
17use crate::response::serialize_response;
18use crate::types::CallerKind;
19
20// Use wee_alloc as the global allocator when targeting WASM to minimise
21// binary size.
22#[cfg(target_arch = "wasm32")]
23#[global_allocator]
24static ALLOC: wee_alloc::WeeAlloc = wee_alloc::WeeAlloc::INIT;
25
26// Store the caller detected during `parse` so that `serialize` can use it
27// without requiring callers to pass it explicitly.
28thread_local! {
29    static LAST_CALLER: RefCell<CallerKind> = RefCell::new(CallerKind::Unknown);
30}
31
32// ---------------------------------------------------------------------------
33// Memory helpers
34// ---------------------------------------------------------------------------
35
36/// Allocate `len` bytes and return a raw pointer.
37///
38/// The allocation is managed by Rust's allocator; the caller is responsible
39/// for calling `dealloc` with the same pointer and length.
40#[no_mangle]
41pub unsafe extern "C" fn alloc(len: usize) -> *mut u8 {
42    let mut buf: Vec<u8> = Vec::with_capacity(len);
43    // Safety: we immediately hand ownership to the caller; they will write
44    // into this memory before passing it back.
45    buf.set_len(len);
46    let ptr = buf.as_mut_ptr();
47    std::mem::forget(buf);
48    ptr
49}
50
51/// Free memory previously allocated by `alloc` or returned by `parse`/`serialize`.
52#[no_mangle]
53pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) {
54    // Reconstruct the Vec so Rust can drop it.
55    let _ = Vec::from_raw_parts(ptr, len, len);
56}
57
58// ---------------------------------------------------------------------------
59// Core exports
60// ---------------------------------------------------------------------------
61
62/// Parse raw JSON bytes into a normalised `HookEvent` and return it as
63/// length-prefixed JSON.
64///
65/// Side-effect: stores the detected `CallerKind` in the thread-local so that
66/// the subsequent `serialize` call can use it.
67#[no_mangle]
68pub unsafe extern "C" fn parse(ptr: *const u8, len: usize) -> *mut u8 {
69    let input = std::slice::from_raw_parts(ptr, len);
70
71    let result: Vec<u8> = match parse_event(input) {
72        Ok(event) => {
73            // Persist the caller for subsequent serialize calls.
74            LAST_CALLER.with(|c| {
75                *c.borrow_mut() = event.caller.clone();
76            });
77            match serde_json::to_vec(&event) {
78                Ok(bytes) => bytes,
79                Err(e) => {
80                    let msg = format!("{{\"error\":\"serialize failed: {e}\"}}");
81                    msg.into_bytes()
82                }
83            }
84        }
85        Err(e) => {
86            let msg = format!("{{\"error\":\"{e}\"}}");
87            msg.into_bytes()
88        }
89    };
90
91    length_prefix_alloc(result)
92}
93
94/// Deserialise a `HookResponse` JSON and re-serialise it in the format
95/// expected by the caller detected during the most recent `parse` call.
96#[no_mangle]
97pub unsafe extern "C" fn serialize(ptr: *const u8, len: usize) -> *mut u8 {
98    let input = std::slice::from_raw_parts(ptr, len);
99
100    // Parse as a generic Value first, then dispatch on "action" to build a
101    // typed HookResponse.  serde untagged deserialization can't safely
102    // disambiguate the variants (ApproveResponse matches everything because it
103    // has no required-unique fields), so we do it manually.
104    let result: Vec<u8> = match serde_json::from_slice::<serde_json::Value>(input) {
105        Ok(val) => {
106            let resp = match val.get("action").and_then(|a| a.as_str()) {
107                Some("block") => {
108                    let msg = val
109                        .get("message")
110                        .and_then(|m| m.as_str())
111                        .unwrap_or("");
112                    crate::types::HookResponse::block(msg)
113                }
114                Some("modify") => {
115                    let input = val
116                        .get("input")
117                        .cloned()
118                        .unwrap_or(serde_json::Value::Object(Default::default()));
119                    crate::types::HookResponse::modify(input)
120                }
121                _ => crate::types::HookResponse::approve(),
122            };
123            let caller = LAST_CALLER.with(|c| c.borrow().clone());
124            let value = serialize_response(&resp, &caller);
125            match serde_json::to_vec(&value) {
126                Ok(bytes) => bytes,
127                Err(e) => format!("{{\"error\":\"serialize failed: {e}\"}}").into_bytes(),
128            }
129        }
130        Err(e) => format!("{{\"error\":\"response parse failed: {e}\"}}").into_bytes(),
131    };
132
133    length_prefix_alloc(result)
134}
135
136// ---------------------------------------------------------------------------
137// Internal helpers
138// ---------------------------------------------------------------------------
139
140/// Prepend a 4-byte little-endian length to `payload`, allocate, and return
141/// a raw pointer to the combined buffer.  The caller owns the memory and must
142/// call `dealloc(ptr, 4 + payload_len)`.
143fn length_prefix_alloc(payload: Vec<u8>) -> *mut u8 {
144    let payload_len = payload.len();
145    let total = 4 + payload_len;
146
147    let mut buf: Vec<u8> = Vec::with_capacity(total);
148    let len_bytes = (payload_len as i32).to_le_bytes();
149    buf.extend_from_slice(&len_bytes);
150    buf.extend_from_slice(&payload);
151
152    debug_assert_eq!(buf.len(), total);
153
154    let ptr = buf.as_mut_ptr();
155    std::mem::forget(buf);
156    ptr
157}
158
159// ---------------------------------------------------------------------------
160// Tests
161// ---------------------------------------------------------------------------
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166
167    // Helper: read a length-prefixed buffer returned by parse/serialize.
168    // Returns the payload bytes and the total buffer length (4 + payload).
169    unsafe fn read_length_prefixed(ptr: *mut u8) -> (Vec<u8>, usize) {
170        // Read 4-byte LE length header.
171        let len_bytes: [u8; 4] = std::slice::from_raw_parts(ptr, 4)
172            .try_into()
173            .expect("slice to array");
174        let payload_len = i32::from_le_bytes(len_bytes) as usize;
175        let total = 4 + payload_len;
176
177        // Copy out the payload.
178        let payload = std::slice::from_raw_parts(ptr.add(4), payload_len).to_vec();
179        (payload, total)
180    }
181
182    // ---------------------------------------------------------------------------
183    // alloc / dealloc
184    // ---------------------------------------------------------------------------
185
186    #[test]
187    fn alloc_zero_does_not_crash() {
188        unsafe {
189            let ptr = alloc(0);
190            // Deallocating a zero-length allocation; ptr may be dangling/null but
191            // Vec::from_raw_parts(ptr, 0, 0) is defined to drop nothing.
192            dealloc(ptr, 0);
193        }
194    }
195
196    #[test]
197    fn alloc_returns_non_null_for_nonzero_len() {
198        unsafe {
199            let ptr = alloc(64);
200            assert!(!ptr.is_null());
201            dealloc(ptr, 64);
202        }
203    }
204
205    #[test]
206    fn dealloc_of_alloc_does_not_crash() {
207        unsafe {
208            let ptr = alloc(128);
209            assert!(!ptr.is_null());
210            // Write something to confirm the memory is usable.
211            std::ptr::write_bytes(ptr, 0xAB, 128);
212            dealloc(ptr, 128);
213        }
214    }
215
216    // ---------------------------------------------------------------------------
217    // parse
218    // ---------------------------------------------------------------------------
219
220    #[test]
221    fn parse_valid_claude_code_json_returns_tool_before() {
222        let json =
223            br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s1"}"#;
224        unsafe {
225            let input_ptr = alloc(json.len());
226            std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
227
228            let out_ptr = parse(input_ptr as *const u8, json.len());
229            assert!(!out_ptr.is_null());
230
231            let (payload, total) = read_length_prefixed(out_ptr);
232            let s = String::from_utf8(payload).expect("valid utf8");
233            assert!(s.contains("tool:before"), "expected 'tool:before' in: {s}");
234
235            dealloc(input_ptr, json.len());
236            dealloc(out_ptr, total);
237        }
238    }
239
240    #[test]
241    fn parse_invalid_json_returns_error_payload() {
242        let bad = b"this is not json";
243        unsafe {
244            let input_ptr = alloc(bad.len());
245            std::ptr::copy_nonoverlapping(bad.as_ptr(), input_ptr, bad.len());
246
247            let out_ptr = parse(input_ptr as *const u8, bad.len());
248            assert!(!out_ptr.is_null());
249
250            let (payload, total) = read_length_prefixed(out_ptr);
251            let s = String::from_utf8(payload).expect("valid utf8");
252            assert!(s.contains("error"), "expected 'error' in: {s}");
253
254            dealloc(input_ptr, bad.len());
255            dealloc(out_ptr, total);
256        }
257    }
258
259    // ---------------------------------------------------------------------------
260    // serialize
261    // ---------------------------------------------------------------------------
262
263    #[test]
264    fn serialize_approve_returns_length_prefixed_json() {
265        let json = br#"{"action":"approve"}"#;
266        unsafe {
267            let input_ptr = alloc(json.len());
268            std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
269
270            let out_ptr = serialize(input_ptr as *const u8, json.len());
271            assert!(!out_ptr.is_null());
272
273            let (payload, total) = read_length_prefixed(out_ptr);
274            let s = String::from_utf8(payload).expect("valid utf8");
275            // Should be valid JSON and not contain "error".
276            let parsed: serde_json::Value =
277                serde_json::from_str(&s).expect("should be valid JSON");
278            assert!(!s.contains("error"), "unexpected error in: {s}");
279            // Result is a JSON object.
280            assert!(parsed.is_object());
281
282            dealloc(input_ptr, json.len());
283            dealloc(out_ptr, total);
284        }
285    }
286
287    #[test]
288    fn serialize_block_returns_length_prefixed_json_with_block_content() {
289        // First parse something so LAST_CALLER is set (ClaudeCode).
290        let event_json =
291            br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s2"}"#;
292        unsafe {
293            let ep = alloc(event_json.len());
294            std::ptr::copy_nonoverlapping(event_json.as_ptr(), ep, event_json.len());
295            let ep_out = parse(ep as *const u8, event_json.len());
296            let (_, ep_total) = read_length_prefixed(ep_out);
297            dealloc(ep, event_json.len());
298            dealloc(ep_out, ep_total);
299        }
300
301        let json = br#"{"action":"block","message":"dangerous command"}"#;
302        unsafe {
303            let input_ptr = alloc(json.len());
304            std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
305
306            let out_ptr = serialize(input_ptr as *const u8, json.len());
307            assert!(!out_ptr.is_null());
308
309            let (payload, total) = read_length_prefixed(out_ptr);
310            let s = String::from_utf8(payload).expect("valid utf8");
311            // Claude Code block format should have "block" in the JSON.
312            assert!(s.contains("block"), "expected 'block' in: {s}");
313
314            dealloc(input_ptr, json.len());
315            dealloc(out_ptr, total);
316        }
317    }
318
319    #[test]
320    fn serialize_modify_returns_length_prefixed_json() {
321        // Ensure LAST_CALLER is set via parse.
322        let event_json =
323            br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s3"}"#;
324        unsafe {
325            let ep = alloc(event_json.len());
326            std::ptr::copy_nonoverlapping(event_json.as_ptr(), ep, event_json.len());
327            let ep_out = parse(ep as *const u8, event_json.len());
328            let (_, ep_total) = read_length_prefixed(ep_out);
329            dealloc(ep, event_json.len());
330            dealloc(ep_out, ep_total);
331        }
332
333        let json = br#"{"action":"modify","input":{"command":"echo safe"}}"#;
334        unsafe {
335            let input_ptr = alloc(json.len());
336            std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
337
338            let out_ptr = serialize(input_ptr as *const u8, json.len());
339            assert!(!out_ptr.is_null());
340
341            let (payload, total) = read_length_prefixed(out_ptr);
342            let s = String::from_utf8(payload).expect("valid utf8");
343            let parsed: serde_json::Value =
344                serde_json::from_str(&s).expect("should be valid JSON");
345            assert!(parsed.is_object());
346            assert!(!s.contains("error"), "unexpected error in: {s}");
347
348            dealloc(input_ptr, json.len());
349            dealloc(out_ptr, total);
350        }
351    }
352
353    #[test]
354    fn serialize_invalid_json_returns_error_payload() {
355        let bad = b"not json at all!";
356        unsafe {
357            let input_ptr = alloc(bad.len());
358            std::ptr::copy_nonoverlapping(bad.as_ptr(), input_ptr, bad.len());
359
360            let out_ptr = serialize(input_ptr as *const u8, bad.len());
361            assert!(!out_ptr.is_null());
362
363            let (payload, total) = read_length_prefixed(out_ptr);
364            let s = String::from_utf8(payload).expect("valid utf8");
365            assert!(s.contains("error"), "expected 'error' in: {s}");
366
367            dealloc(input_ptr, bad.len());
368            dealloc(out_ptr, total);
369        }
370    }
371
372    // ---------------------------------------------------------------------------
373    // length_prefix_alloc (private helper exercised indirectly above, but also
374    // tested directly via the public API round-trip)
375    // ---------------------------------------------------------------------------
376
377    #[test]
378    fn length_prefix_alloc_encodes_correct_length() {
379        let payload = b"hello world";
380        let payload_vec = payload.to_vec();
381        let payload_len = payload_vec.len();
382
383        unsafe {
384            let ptr = length_prefix_alloc(payload_vec);
385            assert!(!ptr.is_null());
386
387            // Read back the 4-byte LE header.
388            let len_bytes: [u8; 4] =
389                std::slice::from_raw_parts(ptr, 4).try_into().unwrap();
390            let decoded_len = i32::from_le_bytes(len_bytes) as usize;
391            assert_eq!(decoded_len, payload_len);
392
393            // Verify payload bytes.
394            let actual_payload = std::slice::from_raw_parts(ptr.add(4), payload_len);
395            assert_eq!(actual_payload, payload);
396
397            dealloc(ptr, 4 + payload_len);
398        }
399    }
400}