diff options
| author | Mason Reed <mason@vector35.com> | 2024-08-01 14:45:39 -0400 |
|---|---|---|
| committer | Mason Reed <mason@vector35.com> | 2024-12-14 13:13:37 -0500 |
| commit | 121c16592476754800b00a3b51595f4799944d04 (patch) | |
| tree | b1f9556704ff169bca8a66be0a0faa448e5e9ecf /rust/src | |
| parent | 487fa4b240ae2c2d5a712dd6e00071f8de57d0fc (diff) | |
Pass length to free register list callback
Allows language bindings like rust to free register lists sanely
Diffstat (limited to 'rust/src')
| -rw-r--r-- | rust/src/architecture.rs | 134 | ||||
| -rw-r--r-- | rust/src/callingconvention.rs | 69 |
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 _ }) } |
