summaryrefslogtreecommitdiff
path: root/rust
diff options
context:
space:
mode:
authorMason Reed <mason@vector35.com>2024-08-01 14:45:39 -0400
committerMason Reed <mason@vector35.com>2024-12-14 13:13:37 -0500
commit121c16592476754800b00a3b51595f4799944d04 (patch)
treeb1f9556704ff169bca8a66be0a0faa448e5e9ecf /rust
parent487fa4b240ae2c2d5a712dd6e00071f8de57d0fc (diff)
Pass length to free register list callback
Allows language bindings like rust to free register lists sanely
Diffstat (limited to 'rust')
-rw-r--r--rust/src/architecture.rs134
-rw-r--r--rust/src/callingconvention.rs69
2 files changed, 118 insertions, 85 deletions
diff --git a/rust/src/architecture.rs b/rust/src/architecture.rs
index b5a1cfdb..e91925b7 100644
--- a/rust/src/architecture.rs
+++ b/rust/src/architecture.rs
@@ -1941,34 +1941,19 @@ where
None => BnString::new("invalid_flag_group").into_raw(),
}
}
-
- fn alloc_register_list<I: Iterator<Item = u32> + ExactSizeIterator>(
- items: I,
- count: &mut usize,
- ) -> *mut u32 {
- let len = items.len();
- *count = len;
-
- if len == 0 {
- ptr::null_mut()
- } else {
- let mut res: Box<[_]> = [len as u32].into_iter().chain(items).collect();
-
- let raw = res.as_mut_ptr();
- mem::forget(res);
-
- unsafe { raw.offset(1) }
- }
- }
-
+
extern "C" fn cb_registers_full_width<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let regs = custom_arch.registers_full_width();
+ let mut regs = custom_arch.registers_full_width();
- alloc_register_list(regs.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
}
extern "C" fn cb_registers_all<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -1976,9 +1961,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let regs = custom_arch.registers_all();
+ let mut regs = custom_arch.registers_all();
- alloc_register_list(regs.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
}
extern "C" fn cb_registers_global<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -1986,9 +1975,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let regs = custom_arch.registers_global();
+ let mut regs = custom_arch.registers_global();
- alloc_register_list(regs.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
}
extern "C" fn cb_registers_system<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -1996,9 +1989,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let regs = custom_arch.registers_system();
+ let mut regs = custom_arch.registers_system();
- alloc_register_list(regs.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
}
extern "C" fn cb_flags<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -2006,9 +2003,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let flags = custom_arch.flags();
+ let mut flags = custom_arch.flags();
- alloc_register_list(flags.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flags.len() };
+ let regs_ptr = flags.as_mut_ptr();
+ mem::forget(flags);
+ regs_ptr as *mut _
}
extern "C" fn cb_flag_write_types<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -2016,9 +2017,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let flag_writes = custom_arch.flag_write_types();
+ let mut flag_writes = custom_arch.flag_write_types();
- alloc_register_list(flag_writes.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flag_writes.len() };
+ let regs_ptr = flag_writes.as_mut_ptr();
+ mem::forget(flag_writes);
+ regs_ptr as *mut _
}
extern "C" fn cb_semantic_flag_classes<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -2026,9 +2031,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let flag_classes = custom_arch.flag_classes();
+ let mut flag_classes = custom_arch.flag_classes();
- alloc_register_list(flag_classes.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flag_classes.len() };
+ let regs_ptr = flag_classes.as_mut_ptr();
+ mem::forget(flag_classes);
+ regs_ptr as *mut _
}
extern "C" fn cb_semantic_flag_groups<A>(ctxt: *mut c_void, count: *mut usize) -> *mut u32
@@ -2036,9 +2045,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let flag_groups = custom_arch.flag_groups();
+ let mut flag_groups = custom_arch.flag_groups();
- alloc_register_list(flag_groups.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flag_groups.len() };
+ let regs_ptr = flag_groups.as_mut_ptr();
+ mem::forget(flag_groups);
+ regs_ptr as *mut _
}
extern "C" fn cb_flag_role<A>(ctxt: *mut c_void, flag: u32, class: u32) -> BNFlagRole
@@ -2068,9 +2081,13 @@ where
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
let class = custom_arch.flag_class_from_id(class);
- let flags = custom_arch.flags_required_for_flag_condition(cond, class);
+ let mut flags = custom_arch.flags_required_for_flag_condition(cond, class);
- alloc_register_list(flags.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flags.len() };
+ let regs_ptr = flags.as_mut_ptr();
+ mem::forget(flags);
+ regs_ptr as *mut _
}
extern "C" fn cb_flags_required_for_semantic_flag_group<A>(
@@ -2084,8 +2101,13 @@ where
let custom_arch = unsafe { &*(ctxt as *mut A) };
if let Some(group) = custom_arch.flag_group_from_id(group) {
- let flags = group.flags_required();
- alloc_register_list(flags.iter().map(|r| r.id()), unsafe { &mut *count })
+ let mut flags = group.flags_required();
+
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = flags.len() };
+ let regs_ptr = flags.as_mut_ptr();
+ mem::forget(flags);
+ regs_ptr as *mut _
} else {
unsafe {
*count = 0;
@@ -2153,8 +2175,13 @@ where
let custom_arch = unsafe { &*(ctxt as *mut A) };
if let Some(write_type) = custom_arch.flag_write_from_id(write_type) {
- let written = write_type.flags_written();
- alloc_register_list(written.iter().map(|f| f.id()), unsafe { &mut *count })
+ let mut written = write_type.flags_written();
+
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = written.len() };
+ let regs_ptr = written.as_mut_ptr();
+ mem::forget(written);
+ regs_ptr as *mut _
} else {
unsafe {
*count = 0;
@@ -2285,15 +2312,13 @@ where
lifter.unimplemented().expr_idx
}
- extern "C" fn cb_free_register_list(_ctxt: *mut c_void, regs: *mut u32) {
+ extern "C" fn cb_free_register_list(_ctxt: *mut c_void, regs: *mut u32, count: usize) {
if regs.is_null() {
return;
}
unsafe {
- let actual_start = regs.offset(-1);
- let len = *actual_start + 1;
- let regs_ptr = ptr::slice_from_raw_parts_mut(actual_start, len.try_into().unwrap());
+ let regs_ptr = ptr::slice_from_raw_parts_mut(regs, count);
let _regs = Box::from_raw(regs_ptr);
}
}
@@ -2362,9 +2387,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let regs = custom_arch.register_stacks();
+ let mut regs = custom_arch.register_stacks();
- alloc_register_list(regs.iter().map(|r| r.id()), unsafe { &mut *count })
+ // SAFETY: Passed in to be written
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
}
extern "C" fn cb_reg_stack_info<A>(
@@ -2420,8 +2449,13 @@ where
A: 'static + Architecture<Handle = CustomArchitectureHandle<A>> + Send + Sync,
{
let custom_arch = unsafe { &*(ctxt as *mut A) };
- let intrinsics = custom_arch.intrinsics();
- alloc_register_list(intrinsics.iter().map(|i| i.id()), unsafe { &mut *count })
+ let mut intrinsics = custom_arch.intrinsics();
+
+ // SAFETY: Passed in to be written
+ unsafe { *count = intrinsics.len() };
+ let regs_ptr = intrinsics.as_mut_ptr();
+ mem::forget(intrinsics);
+ regs_ptr as *mut _
}
extern "C" fn cb_intrinsic_inputs<A>(
diff --git a/rust/src/callingconvention.rs b/rust/src/callingconvention.rs
index a009b439..915bdcc2 100644
--- a/rust/src/callingconvention.rs
+++ b/rust/src/callingconvention.rs
@@ -81,34 +81,13 @@ where
})
}
- fn alloc_register_list<I: Iterator<Item = u32> + ExactSizeIterator>(
- items: I,
- count: &mut usize,
- ) -> *mut u32 {
- let len = items.len();
- *count = len;
-
- if len == 0 {
- return ptr::null_mut();
- }
-
- let res: Box<[_]> = [len as u32].into_iter().chain(items).collect();
- debug_assert!(res.len() == len + 1);
-
- // it's free on the function below: `cb_free_register_list`
- let raw = Box::leak(res);
- &mut raw[1]
- }
-
- extern "C" fn cb_free_register_list(_ctxt: *mut c_void, regs: *mut u32) {
+ extern "C" fn cb_free_register_list(_ctxt: *mut c_void, regs: *mut u32, count: usize) {
ffi_wrap!("CallingConvention::free_register_list", unsafe {
if regs.is_null() {
return;
}
-
- let actual_start = regs.offset(-1);
- let len = (*actual_start) + 1;
- let _regs = Box::from_raw(ptr::slice_from_raw_parts_mut(actual_start, len as usize));
+
+ let _regs = Box::from_raw(ptr::slice_from_raw_parts_mut(regs, count));
})
}
@@ -118,9 +97,13 @@ where
{
ffi_wrap!("CallingConvention::caller_saved_registers", unsafe {
let ctxt = &*(ctxt as *mut CustomCallingConventionContext<C>);
- let regs = ctxt.cc.caller_saved_registers();
+ let mut regs = ctxt.cc.caller_saved_registers();
- alloc_register_list(regs.iter().map(|r| r.id()), &mut *count)
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
})
}
@@ -130,9 +113,13 @@ where
{
ffi_wrap!("CallingConvention::callee_saved_registers", unsafe {
let ctxt = &*(ctxt as *mut CustomCallingConventionContext<C>);
- let regs = ctxt.cc.callee_saved_registers();
-
- alloc_register_list(regs.iter().map(|r| r.id()), &mut *count)
+ let mut regs = ctxt.cc.callee_saved_registers();
+
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
})
}
@@ -142,9 +129,13 @@ where
{
ffi_wrap!("CallingConvention::int_arg_registers", unsafe {
let ctxt = &*(ctxt as *mut CustomCallingConventionContext<C>);
- let regs = ctxt.cc.int_arg_registers();
+ let mut regs = ctxt.cc.int_arg_registers();
- alloc_register_list(regs.iter().map(|r| r.id()), &mut *count)
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
})
}
@@ -154,9 +145,13 @@ where
{
ffi_wrap!("CallingConvention::float_arg_registers", unsafe {
let ctxt = &*(ctxt as *mut CustomCallingConventionContext<C>);
- let regs = ctxt.cc.float_arg_registers();
+ let mut regs = ctxt.cc.float_arg_registers();
- alloc_register_list(regs.iter().map(|r| r.id()), &mut *count)
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
})
}
@@ -272,9 +267,13 @@ where
{
ffi_wrap!("CallingConvention::implicitly_defined_registers", unsafe {
let ctxt = &*(ctxt as *mut CustomCallingConventionContext<C>);
- let regs = ctxt.cc.implicitly_defined_registers();
+ let mut regs = ctxt.cc.implicitly_defined_registers();
- alloc_register_list(regs.iter().map(|r| r.id()), &mut *count)
+ // SAFETY: `count` is an out parameter
+ unsafe { *count = regs.len() };
+ let regs_ptr = regs.as_mut_ptr();
+ mem::forget(regs);
+ regs_ptr as *mut _
})
}