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