From 534627a04c77aa6791be25ad04deefa9b37034f5 Mon Sep 17 00:00:00 2001 From: Mason Reed Date: Sun, 3 Aug 2025 16:59:02 -0400 Subject: [Rust] Take download callbacks by reference to avoid boxing --- rust/src/download_provider.rs | 18 +++++++----------- rust/tests/download_provider.rs | 34 ++++++++++++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 11 deletions(-) create mode 100644 rust/tests/download_provider.rs (limited to 'rust') diff --git a/rust/src/download_provider.rs b/rust/src/download_provider.rs index 9fd803a7..b0c25b42 100644 --- a/rust/src/download_provider.rs +++ b/rust/src/download_provider.rs @@ -134,14 +134,13 @@ impl DownloadInstance { pub fn perform_request( &mut self, url: &str, - callbacks: DownloadInstanceOutputCallbacks, + callbacks: &DownloadInstanceOutputCallbacks, ) -> Result<(), String> { - let callbacks = Box::into_raw(Box::new(callbacks)); let mut cbs = BNDownloadInstanceOutputCallbacks { writeCallback: Some(Self::o_write_callback), - writeContext: callbacks as *mut c_void, + writeContext: callbacks as *const _ as *mut c_void, progressCallback: Some(Self::o_progress_callback), - progressContext: callbacks as *mut c_void, + progressContext: callbacks as *const _ as *mut c_void, }; let url_raw = url.to_cstr(); @@ -153,8 +152,6 @@ impl DownloadInstance { ) }; - // Drop it - unsafe { drop(Box::from_raw(callbacks)) }; if result < 0 { Err(self.get_error()) } else { @@ -206,7 +203,7 @@ impl DownloadInstance { method: &str, url: &str, headers: I, - callbacks: DownloadInstanceInputOutputCallbacks, + callbacks: &DownloadInstanceInputOutputCallbacks, ) -> Result where I: IntoIterator, @@ -226,14 +223,13 @@ impl DownloadInstance { header_value_ptrs.push(value.as_ptr()); } - let callbacks = Box::into_raw(Box::new(callbacks)); let mut cbs = BNDownloadInstanceInputOutputCallbacks { readCallback: Some(Self::i_read_callback), - readContext: callbacks as *mut c_void, + readContext: callbacks as *const _ as *mut c_void, writeCallback: Some(Self::i_write_callback), - writeContext: callbacks as *mut c_void, + writeContext: callbacks as *const _ as *mut c_void, progressCallback: Some(Self::i_progress_callback), - progressContext: callbacks as *mut c_void, + progressContext: callbacks as *const _ as *mut c_void, }; let mut response: *mut BNDownloadInstanceResponse = null_mut(); diff --git a/rust/tests/download_provider.rs b/rust/tests/download_provider.rs new file mode 100644 index 00000000..f3b62263 --- /dev/null +++ b/rust/tests/download_provider.rs @@ -0,0 +1,34 @@ +use binaryninja::download_provider::{DownloadInstanceInputOutputCallbacks, DownloadProvider}; +use binaryninja::headless::Session; +use std::sync::mpsc; + +#[test] +fn test_download_provider() { + let _session = Session::new().expect("Failed to initialize session"); + let provider = DownloadProvider::try_default().expect("Couldn't get default download provider"); + let mut inst = provider + .create_instance() + .expect("Couldn't create download instance"); + let (tx, rx) = mpsc::channel(); + let write = move |data: &[u8]| -> usize { + tx.send(data.to_vec()).expect("Couldn't send data"); + data.len() + }; + let result = inst + .perform_custom_request( + "GET", + "http://httpbin.org/get", + vec![], + &DownloadInstanceInputOutputCallbacks { + read: None, + write: Some(Box::new(write)), + progress: None, + }, + ) + .expect("Couldn't perform custom request"); + assert_eq!(result.status_code, 200); + let written = rx.recv().expect("Couldn't receive data"); + let written_str = String::from_utf8(written).expect("Couldn't convert data to string"); + println!("{}", written_str); + assert!(written_str.contains("httpbin.org/get")); +} -- cgit v1.3.1