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_with_event;
18use crate::types::{CallerKind, HookEventEvent};
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 and event type detected during `parse` so that `serialize`
27// can use them without requiring callers to pass them explicitly. Mirrors the
28// LAST_CALLER/LAST_EVENT pair in lib.rs's native `read_from`/`respond_to`
29// path, which this WASM ABI must match to pick the correct Claude Code block
30// format (see serialize_response_with_event).
31thread_local! {
32    static LAST_CALLER: RefCell<CallerKind> = const { RefCell::new(CallerKind::Unknown) };
33    static LAST_EVENT: RefCell<Option<HookEventEvent>> = const { RefCell::new(None) };
34}
35
36// ---------------------------------------------------------------------------
37// Memory helpers
38// ---------------------------------------------------------------------------
39
40/// Allocate `len` bytes and return a raw pointer.
41///
42/// The allocation is managed by Rust's allocator; the caller is responsible
43/// for calling `dealloc` with the same pointer and length.
44///
45/// # Safety
46///
47/// The caller must write exactly `len` bytes before passing the pointer back
48/// to `parse` or `serialize`, and must call `dealloc(ptr, len)` exactly once
49/// when done.
50#[no_mangle]
51pub unsafe extern "C" fn alloc(len: usize) -> *mut u8 {
52    let mut buf: Vec<u8> = Vec::with_capacity(len);
53    let ptr = buf.as_mut_ptr();
54    std::mem::forget(buf);
55    ptr
56}
57
58/// Free memory previously allocated by `alloc` or returned by `parse`/`serialize`.
59///
60/// # Safety
61///
62/// `ptr` must have been returned by `alloc`, `parse`, or `serialize`, and
63/// `len` must be the exact byte count that was originally allocated.  Must
64/// not be called more than once for the same pointer.
65#[no_mangle]
66pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) {
67    // Reconstruct the Vec so Rust can drop it.
68    let _ = Vec::from_raw_parts(ptr, len, len);
69}
70
71// ---------------------------------------------------------------------------
72// Core exports
73// ---------------------------------------------------------------------------
74
75/// Parse raw JSON bytes into a normalised `HookEvent` and return it as
76/// length-prefixed JSON.
77///
78/// Side-effect: stores the detected `CallerKind` in the thread-local so that
79/// the subsequent `serialize` call can use it.
80///
81/// # Safety
82///
83/// `ptr` must point to `len` consecutive readable bytes valid for the
84/// duration of this call.  The returned pointer must be freed with
85/// `dealloc(ptr, 4 + payload_len)` where `payload_len` is the LE-i32 at the
86/// first four bytes of the returned buffer.
87#[no_mangle]
88pub unsafe extern "C" fn parse(ptr: *const u8, len: usize) -> *mut u8 {
89    let input = std::slice::from_raw_parts(ptr, len);
90
91    let result: Vec<u8> = match parse_event(input) {
92        Ok(event) => {
93            // Persist the caller and event type for the subsequent serialize call.
94            LAST_CALLER.with(|c| {
95                *c.borrow_mut() = event.caller;
96            });
97            LAST_EVENT.with(|e| {
98                *e.borrow_mut() = Some(event.event);
99            });
100            match serde_json::to_vec(&event) {
101                Ok(bytes) => bytes,
102                Err(e) => {
103                    let msg = format!("{{\"error\":\"serialize failed: {e}\"}}");
104                    msg.into_bytes()
105                }
106            }
107        }
108        Err(e) => {
109            let msg = format!("{{\"error\":\"{e}\"}}");
110            msg.into_bytes()
111        }
112    };
113
114    length_prefix_alloc(result)
115}
116
117/// Deserialise a `HookResponse` JSON and re-serialise it in the format
118/// expected by the caller detected during the most recent `parse` call.
119///
120/// # Safety
121///
122/// `ptr` must point to `len` consecutive readable bytes valid for the
123/// duration of this call.  The returned pointer must be freed with
124/// `dealloc(ptr, 4 + payload_len)` where `payload_len` is the LE-i32 at the
125/// first four bytes of the returned buffer.
126#[no_mangle]
127pub unsafe extern "C" fn serialize(ptr: *const u8, len: usize) -> *mut u8 {
128    let input = std::slice::from_raw_parts(ptr, len);
129
130    // Parse as a generic Value first, then dispatch on "action" to build a
131    // typed HookResponse.  serde untagged deserialization can't safely
132    // disambiguate the variants (ApproveResponse matches everything because it
133    // has no required-unique fields), so we do it manually.
134    let result: Vec<u8> = match serde_json::from_slice::<serde_json::Value>(input) {
135        Ok(val) => {
136            let resp = match val.get("action").and_then(|a| a.as_str()) {
137                Some("block") => {
138                    let msg = val.get("message").and_then(|m| m.as_str()).unwrap_or("");
139                    crate::types::HookResponse::block(msg)
140                }
141                Some("modify") => {
142                    let input = val
143                        .get("input")
144                        .cloned()
145                        .unwrap_or(serde_json::Value::Object(Default::default()));
146                    crate::types::HookResponse::modify(input)
147                }
148                Some("context") => {
149                    let context = val.get("context").and_then(|c| c.as_str()).unwrap_or("");
150                    crate::types::HookResponse::context(context)
151                }
152                _ => crate::types::HookResponse::approve(),
153            };
154            let caller = LAST_CALLER.with(|c| *c.borrow());
155            let event = LAST_EVENT.with(|e| *e.borrow());
156            let value = serialize_response_with_event(&resp, caller, event);
157            match serde_json::to_vec(&value) {
158                Ok(bytes) => bytes,
159                Err(e) => format!("{{\"error\":\"serialize failed: {e}\"}}").into_bytes(),
160            }
161        }
162        Err(e) => format!("{{\"error\":\"response parse failed: {e}\"}}").into_bytes(),
163    };
164
165    length_prefix_alloc(result)
166}
167
168// ---------------------------------------------------------------------------
169// Internal helpers
170// ---------------------------------------------------------------------------
171
172/// Prepend a 4-byte little-endian length to `payload`, allocate, and return
173/// a raw pointer to the combined buffer.  The caller owns the memory and must
174/// call `dealloc(ptr, 4 + payload_len)`.
175fn length_prefix_alloc(payload: Vec<u8>) -> *mut u8 {
176    let payload_len = payload.len();
177    let total = 4 + payload_len;
178
179    let mut buf: Vec<u8> = Vec::with_capacity(total);
180    let len_bytes = (payload_len as i32).to_le_bytes();
181    buf.extend_from_slice(&len_bytes);
182    buf.extend_from_slice(&payload);
183
184    debug_assert_eq!(buf.len(), total);
185
186    let ptr = buf.as_mut_ptr();
187    std::mem::forget(buf);
188    ptr
189}
190
191// ---------------------------------------------------------------------------
192// Tests
193// ---------------------------------------------------------------------------
194
195#[cfg(test)]
196#[path = "wasm_tests.rs"]
197mod tests;