use std::cell::RefCell;
use crate::parse::parse_event;
use crate::response::serialize_response;
use crate::types::CallerKind;
#[cfg(target_arch = "wasm32")]
#[global_allocator]
static ALLOC: wee_alloc::WeeAlloc = wee_alloc::WeeAlloc::INIT;
thread_local! {
static LAST_CALLER: RefCell<CallerKind> = RefCell::new(CallerKind::Unknown);
}
#[no_mangle]
pub unsafe extern "C" fn alloc(len: usize) -> *mut u8 {
let mut buf: Vec<u8> = Vec::with_capacity(len);
buf.set_len(len);
let ptr = buf.as_mut_ptr();
std::mem::forget(buf);
ptr
}
#[no_mangle]
pub unsafe extern "C" fn dealloc(ptr: *mut u8, len: usize) {
let _ = Vec::from_raw_parts(ptr, len, len);
}
#[no_mangle]
pub unsafe extern "C" fn parse(ptr: *const u8, len: usize) -> *mut u8 {
let input = std::slice::from_raw_parts(ptr, len);
let result: Vec<u8> = match parse_event(input) {
Ok(event) => {
LAST_CALLER.with(|c| {
*c.borrow_mut() = event.caller.clone();
});
match serde_json::to_vec(&event) {
Ok(bytes) => bytes,
Err(e) => {
let msg = format!("{{\"error\":\"serialize failed: {e}\"}}");
msg.into_bytes()
}
}
}
Err(e) => {
let msg = format!("{{\"error\":\"{e}\"}}");
msg.into_bytes()
}
};
length_prefix_alloc(result)
}
#[no_mangle]
pub unsafe extern "C" fn serialize(ptr: *const u8, len: usize) -> *mut u8 {
let input = std::slice::from_raw_parts(ptr, len);
let result: Vec<u8> = match serde_json::from_slice::<serde_json::Value>(input) {
Ok(val) => {
let resp = match val.get("action").and_then(|a| a.as_str()) {
Some("block") => {
let msg = val
.get("message")
.and_then(|m| m.as_str())
.unwrap_or("");
crate::types::HookResponse::block(msg)
}
Some("modify") => {
let input = val
.get("input")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
crate::types::HookResponse::modify(input)
}
_ => crate::types::HookResponse::approve(),
};
let caller = LAST_CALLER.with(|c| c.borrow().clone());
let value = serialize_response(&resp, &caller);
match serde_json::to_vec(&value) {
Ok(bytes) => bytes,
Err(e) => format!("{{\"error\":\"serialize failed: {e}\"}}").into_bytes(),
}
}
Err(e) => format!("{{\"error\":\"response parse failed: {e}\"}}").into_bytes(),
};
length_prefix_alloc(result)
}
fn length_prefix_alloc(payload: Vec<u8>) -> *mut u8 {
let payload_len = payload.len();
let total = 4 + payload_len;
let mut buf: Vec<u8> = Vec::with_capacity(total);
let len_bytes = (payload_len as i32).to_le_bytes();
buf.extend_from_slice(&len_bytes);
buf.extend_from_slice(&payload);
debug_assert_eq!(buf.len(), total);
let ptr = buf.as_mut_ptr();
std::mem::forget(buf);
ptr
}
#[cfg(test)]
mod tests {
use super::*;
unsafe fn read_length_prefixed(ptr: *mut u8) -> (Vec<u8>, usize) {
let len_bytes: [u8; 4] = std::slice::from_raw_parts(ptr, 4)
.try_into()
.expect("slice to array");
let payload_len = i32::from_le_bytes(len_bytes) as usize;
let total = 4 + payload_len;
let payload = std::slice::from_raw_parts(ptr.add(4), payload_len).to_vec();
(payload, total)
}
#[test]
fn alloc_zero_does_not_crash() {
unsafe {
let ptr = alloc(0);
dealloc(ptr, 0);
}
}
#[test]
fn alloc_returns_non_null_for_nonzero_len() {
unsafe {
let ptr = alloc(64);
assert!(!ptr.is_null());
dealloc(ptr, 64);
}
}
#[test]
fn dealloc_of_alloc_does_not_crash() {
unsafe {
let ptr = alloc(128);
assert!(!ptr.is_null());
std::ptr::write_bytes(ptr, 0xAB, 128);
dealloc(ptr, 128);
}
}
#[test]
fn parse_valid_claude_code_json_returns_tool_before() {
let json =
br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s1"}"#;
unsafe {
let input_ptr = alloc(json.len());
std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
let out_ptr = parse(input_ptr as *const u8, json.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
assert!(s.contains("tool:before"), "expected 'tool:before' in: {s}");
dealloc(input_ptr, json.len());
dealloc(out_ptr, total);
}
}
#[test]
fn parse_invalid_json_returns_error_payload() {
let bad = b"this is not json";
unsafe {
let input_ptr = alloc(bad.len());
std::ptr::copy_nonoverlapping(bad.as_ptr(), input_ptr, bad.len());
let out_ptr = parse(input_ptr as *const u8, bad.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
assert!(s.contains("error"), "expected 'error' in: {s}");
dealloc(input_ptr, bad.len());
dealloc(out_ptr, total);
}
}
#[test]
fn serialize_approve_returns_length_prefixed_json() {
let json = br#"{"action":"approve"}"#;
unsafe {
let input_ptr = alloc(json.len());
std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
let out_ptr = serialize(input_ptr as *const u8, json.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
let parsed: serde_json::Value =
serde_json::from_str(&s).expect("should be valid JSON");
assert!(!s.contains("error"), "unexpected error in: {s}");
assert!(parsed.is_object());
dealloc(input_ptr, json.len());
dealloc(out_ptr, total);
}
}
#[test]
fn serialize_block_returns_length_prefixed_json_with_block_content() {
let event_json =
br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s2"}"#;
unsafe {
let ep = alloc(event_json.len());
std::ptr::copy_nonoverlapping(event_json.as_ptr(), ep, event_json.len());
let ep_out = parse(ep as *const u8, event_json.len());
let (_, ep_total) = read_length_prefixed(ep_out);
dealloc(ep, event_json.len());
dealloc(ep_out, ep_total);
}
let json = br#"{"action":"block","message":"dangerous command"}"#;
unsafe {
let input_ptr = alloc(json.len());
std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
let out_ptr = serialize(input_ptr as *const u8, json.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
assert!(s.contains("block"), "expected 'block' in: {s}");
dealloc(input_ptr, json.len());
dealloc(out_ptr, total);
}
}
#[test]
fn serialize_modify_returns_length_prefixed_json() {
let event_json =
br#"{"type":"PreToolUse","tool_name":"Bash","tool_input":{"command":"ls"},"session_id":"s3"}"#;
unsafe {
let ep = alloc(event_json.len());
std::ptr::copy_nonoverlapping(event_json.as_ptr(), ep, event_json.len());
let ep_out = parse(ep as *const u8, event_json.len());
let (_, ep_total) = read_length_prefixed(ep_out);
dealloc(ep, event_json.len());
dealloc(ep_out, ep_total);
}
let json = br#"{"action":"modify","input":{"command":"echo safe"}}"#;
unsafe {
let input_ptr = alloc(json.len());
std::ptr::copy_nonoverlapping(json.as_ptr(), input_ptr, json.len());
let out_ptr = serialize(input_ptr as *const u8, json.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
let parsed: serde_json::Value =
serde_json::from_str(&s).expect("should be valid JSON");
assert!(parsed.is_object());
assert!(!s.contains("error"), "unexpected error in: {s}");
dealloc(input_ptr, json.len());
dealloc(out_ptr, total);
}
}
#[test]
fn serialize_invalid_json_returns_error_payload() {
let bad = b"not json at all!";
unsafe {
let input_ptr = alloc(bad.len());
std::ptr::copy_nonoverlapping(bad.as_ptr(), input_ptr, bad.len());
let out_ptr = serialize(input_ptr as *const u8, bad.len());
assert!(!out_ptr.is_null());
let (payload, total) = read_length_prefixed(out_ptr);
let s = String::from_utf8(payload).expect("valid utf8");
assert!(s.contains("error"), "expected 'error' in: {s}");
dealloc(input_ptr, bad.len());
dealloc(out_ptr, total);
}
}
#[test]
fn length_prefix_alloc_encodes_correct_length() {
let payload = b"hello world";
let payload_vec = payload.to_vec();
let payload_len = payload_vec.len();
unsafe {
let ptr = length_prefix_alloc(payload_vec);
assert!(!ptr.is_null());
let len_bytes: [u8; 4] =
std::slice::from_raw_parts(ptr, 4).try_into().unwrap();
let decoded_len = i32::from_le_bytes(len_bytes) as usize;
assert_eq!(decoded_len, payload_len);
let actual_payload = std::slice::from_raw_parts(ptr.add(4), payload_len);
assert_eq!(actual_payload, payload);
dealloc(ptr, 4 + payload_len);
}
}
}