1use std::cell::RefCell;
15
16use crate::parse::parse_event;
17use crate::response::serialize_response;
18use crate::types::CallerKind;
19
20#[cfg(target_arch = "wasm32")]
23#[global_allocator]
24static ALLOC: wee_alloc::WeeAlloc = wee_alloc::WeeAlloc::INIT;
25
26thread_local! {
29 static LAST_CALLER: RefCell<CallerKind> = const { RefCell::new(CallerKind::Unknown) };
30}
31
32#[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#[no_mangle]
62pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) {
63 let _ = Vec::from_raw_parts(ptr, len, len);
65}
66
67#[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 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#[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 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
156fn 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#[cfg(test)]
184mod tests {
185 use super::*;
186
187 unsafe fn read_length_prefixed(ptr: *mut u8) -> (Vec<u8>, usize) {
190 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 let payload = std::slice::from_raw_parts(ptr.add(4), payload_len).to_vec();
199 (payload, total)
200 }
201
202 #[test]
207 fn alloc_zero_does_not_crash() {
208 unsafe {
209 let ptr = alloc(0);
210 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 std::ptr::write_bytes(ptr, 0xAB, 128);
232 dealloc(ptr, 128);
233 }
234 }
235
236 #[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 #[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 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 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 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 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 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 #[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 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 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}