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> = RefCell::new(CallerKind::Unknown);
30}
31
32#[no_mangle]
41pub unsafe extern "C" fn alloc(len: usize) -> *mut u8 {
42 let mut buf: Vec<u8> = Vec::with_capacity(len);
43 buf.set_len(len);
46 let ptr = buf.as_mut_ptr();
47 std::mem::forget(buf);
48 ptr
49}
50
51#[no_mangle]
53pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) {
54 let _ = Vec::from_raw_parts(ptr, len, len);
56}
57
58#[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 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#[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 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
136fn 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#[cfg(test)]
164mod tests {
165 use super::*;
166
167 unsafe fn read_length_prefixed(ptr: *mut u8) -> (Vec<u8>, usize) {
170 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 let payload = std::slice::from_raw_parts(ptr.add(4), payload_len).to_vec();
179 (payload, total)
180 }
181
182 #[test]
187 fn alloc_zero_does_not_crash() {
188 unsafe {
189 let ptr = alloc(0);
190 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 std::ptr::write_bytes(ptr, 0xAB, 128);
212 dealloc(ptr, 128);
213 }
214 }
215
216 #[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 #[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 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 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 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 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 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 #[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 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 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}