diff options
| author | Mason Reed <mason@vector35.com> | 2025-09-27 17:39:19 -0400 |
|---|---|---|
| committer | Mason Reed <mason@vector35.com> | 2025-10-22 00:36:24 -0400 |
| commit | 31ec051fd683d3747f10a0981b497d90f08d25a3 (patch) | |
| tree | 039f28498656a427297406d90955766981bc0b06 /rust/src/download/instance.rs | |
| parent | 058800a116bf25f1f215b758663366d4695f0e31 (diff) | |
[Rust] Refactor download provider module to allow for custom implementations
Diffstat (limited to 'rust/src/download/instance.rs')
| -rw-r--r-- | rust/src/download/instance.rs | 484 |
1 files changed, 484 insertions, 0 deletions
diff --git a/rust/src/download/instance.rs b/rust/src/download/instance.rs new file mode 100644 index 00000000..7064024b --- /dev/null +++ b/rust/src/download/instance.rs @@ -0,0 +1,484 @@ +use crate::download::DownloadProvider; +use crate::headless::is_shutdown_requested; +use crate::rc::{Ref, RefCountable}; +use crate::string::{strings_to_string_list, BnString, IntoCStr}; +use binaryninjacore_sys::*; +use std::collections::HashMap; +use std::ffi::{c_void, CStr}; +use std::mem::{ManuallyDrop, MaybeUninit}; +use std::os::raw::c_char; +use std::ptr::null_mut; +use std::slice; + +pub trait CustomDownloadInstance: Sized { + fn new_with_provider(provider: DownloadProvider) -> Result<Ref<DownloadInstance>, ()> { + let instance_uninit = MaybeUninit::uninit(); + // SAFETY: Download instance is freed by cb_destroy_instance + let leaked_instance = Box::leak(Box::new(instance_uninit)); + let mut callbacks = BNDownloadInstanceCallbacks { + context: leaked_instance as *mut _ as *mut c_void, + destroyInstance: Some(cb_destroy_instance::<Self>), + performRequest: Some(cb_perform_request::<Self>), + performCustomRequest: Some(cb_perform_custom_request::<Self>), + freeResponse: Some(cb_free_response::<Self>), + }; + let instance_ptr = unsafe { BNInitDownloadInstance(provider.handle, &mut callbacks) }; + // TODO: If possible pass a sensible error back... + let instance_ref = unsafe { DownloadInstance::ref_from_raw(instance_ptr) }; + // We now have the core instance, so we can actually construct the object. + leaked_instance.write(Self::from_core(instance_ref.clone())); + Ok(instance_ref) + } + + /// Construct the object now that the core object has been created. + fn from_core(core: Ref<DownloadInstance>) -> Self; + + /// Get the core object, typically the handle is stored directly on the object. + fn handle(&self) -> Ref<DownloadInstance>; + + /// Send an HTTP request on behalf of the caller. + /// + /// The caller will expect you to inform them of progress via the following: + /// + /// - [DownloadInstance::read_callback] + /// - [DownloadInstance::write_callback] + /// - [DownloadInstance::progress_callback] + fn perform_request(&self, url: &str) -> Result<(), String> { + self.perform_custom_request("GET", url, vec![])?; + Ok(()) + } + + /// Send an HTTP request on behalf of the caller. + /// + /// The caller will expect you to inform them of progress via the following: + /// + /// - [DownloadInstance::read_callback] + /// - [DownloadInstance::write_callback] + /// - [DownloadInstance::progress_callback] + fn perform_custom_request<I>( + &self, + method: &str, + url: &str, + headers: I, + ) -> Result<DownloadResponse, String> + where + I: IntoIterator<Item = (String, String)>; +} + +// TODO: Change this to a trait? +pub struct DownloadInstanceOutputCallbacks { + pub write: Option<Box<dyn FnMut(&[u8]) -> usize>>, + pub progress: Option<Box<dyn FnMut(usize, usize) -> bool>>, +} + +// TODO: Change this to a trait? +pub struct DownloadInstanceInputOutputCallbacks { + pub read: Option<Box<dyn FnMut(&mut [u8]) -> Option<usize>>>, + pub write: Option<Box<dyn FnMut(&[u8]) -> usize>>, + pub progress: Option<Box<dyn FnMut(usize, usize) -> bool>>, +} + +pub struct DownloadResponse { + pub status_code: u16, + pub headers: HashMap<String, String>, +} + +/// A reader for a [`DownloadInstance`]. +pub struct DownloadInstanceReader { + pub instance: Ref<DownloadInstance>, +} + +impl DownloadInstanceReader { + pub fn new(instance: Ref<DownloadInstance>) -> Self { + Self { instance } + } +} + +impl std::io::Read for DownloadInstanceReader { + fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> { + let length = self.instance.read_callback(buf); + if length < 0 { + Err(std::io::Error::new( + std::io::ErrorKind::Interrupted, + "Connection interrupted", + )) + } else { + Ok(length as usize) + } + } +} + +/// A writer for a [`DownloadInstance`]. +pub struct DownloadInstanceWriter { + pub instance: Ref<DownloadInstance>, + /// The expected length of the download. + pub total_length: Option<u64>, + /// The current progress of the download. + pub progress: u64, +} + +impl DownloadInstanceWriter { + pub fn new(instance: Ref<DownloadInstance>, total_length: Option<u64>) -> Self { + Self { + instance, + total_length, + progress: 0, + } + } +} + +impl std::io::Write for DownloadInstanceWriter { + fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> { + let length = self.instance.write_callback(buf); + if is_shutdown_requested() || length == 0 { + Err(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)) + } else { + self.progress += buf.len() as u64; + if self + .instance + .progress_callback(self.progress, self.total_length.unwrap_or(u64::max_value())) + { + Ok(length as usize) + } else { + Err(std::io::Error::from(std::io::ErrorKind::ConnectionAborted)) + } + } + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +pub struct DownloadInstance { + pub(crate) handle: *mut BNDownloadInstance, +} + +impl DownloadInstance { + pub(crate) unsafe fn from_raw(handle: *mut BNDownloadInstance) -> Self { + debug_assert!(!handle.is_null()); + Self { handle } + } + + pub(crate) unsafe fn ref_from_raw(handle: *mut BNDownloadInstance) -> Ref<Self> { + Ref::new(Self::from_raw(handle)) + } + + fn get_error(&self) -> String { + let err: *mut c_char = unsafe { BNGetErrorForDownloadInstance(self.handle) }; + unsafe { BnString::into_string(err) } + } + + /// Sets the error for the instance, any later call to [`DownloadInstance::get_error`] will + /// return the string passed here. + fn set_error(&self, err: &str) { + let err = err.to_cstr(); + unsafe { BNSetErrorForDownloadInstance(self.handle, err.as_ptr()) }; + } + + /// Use inside [`CustomDownloadInstance::perform_custom_request`] to pass data back to the caller. + pub fn write_callback(&self, data: &[u8]) -> u64 { + unsafe { + BNWriteDataForDownloadInstance(self.handle, data.as_ptr() as *mut _, data.len() as u64) + } + } + + /// Use inside [`CustomDownloadInstance::perform_custom_request`] to read data from the caller. + pub fn read_callback(&self, data: &mut [u8]) -> i64 { + unsafe { + BNReadDataForDownloadInstance( + self.handle, + data.as_mut_ptr() as *mut _, + data.len() as u64, + ) + } + } + + /// Use inside [`CustomDownloadInstance::perform_custom_request`] to inform the caller of the request progress. + pub fn progress_callback(&self, progress: u64, total: u64) -> bool { + unsafe { BNNotifyProgressForDownloadInstance(self.handle, progress, total) } + } + + pub fn perform_request( + &mut self, + url: &str, + callbacks: &DownloadInstanceOutputCallbacks, + ) -> Result<(), String> { + let mut cbs = BNDownloadInstanceOutputCallbacks { + writeCallback: Some(cb_write_output), + writeContext: callbacks as *const _ as *mut c_void, + progressCallback: Some(cb_progress_output), + progressContext: callbacks as *const _ as *mut c_void, + }; + + let url_raw = url.to_cstr(); + let result = unsafe { + BNPerformDownloadRequest( + self.handle, + url_raw.as_ptr(), + &mut cbs as *mut BNDownloadInstanceOutputCallbacks, + ) + }; + + if result < 0 { + Err(self.get_error()) + } else { + Ok(()) + } + } + + pub fn perform_custom_request<I>( + &mut self, + method: &str, + url: &str, + headers: I, + callbacks: &DownloadInstanceInputOutputCallbacks, + ) -> Result<DownloadResponse, String> + where + I: IntoIterator<Item = (String, String)>, + { + let mut header_keys = vec![]; + let mut header_values = vec![]; + for (key, value) in headers { + header_keys.push(key.to_cstr()); + header_values.push(value.to_cstr()); + } + + let mut header_key_ptrs = vec![]; + let mut header_value_ptrs = vec![]; + + for (key, value) in header_keys.iter().zip(header_values.iter()) { + header_key_ptrs.push(key.as_ptr()); + header_value_ptrs.push(value.as_ptr()); + } + + let mut cbs = BNDownloadInstanceInputOutputCallbacks { + readCallback: Some(cb_read_input), + readContext: callbacks as *const _ as *mut c_void, + writeCallback: Some(cb_write_input), + writeContext: callbacks as *const _ as *mut c_void, + progressCallback: Some(cb_progress_input), + progressContext: callbacks as *const _ as *mut c_void, + }; + + let mut response: *mut BNDownloadInstanceResponse = null_mut(); + + let method_raw = method.to_cstr(); + let url_raw = url.to_cstr(); + let result = unsafe { + BNPerformCustomRequest( + self.handle, + method_raw.as_ptr(), + url_raw.as_ptr(), + header_key_ptrs.len() as u64, + header_key_ptrs.as_ptr(), + header_value_ptrs.as_ptr(), + &mut response as *mut *mut BNDownloadInstanceResponse, + &mut cbs as *mut BNDownloadInstanceInputOutputCallbacks, + ) + }; + + if result < 0 { + unsafe { BNFreeDownloadInstanceResponse(response) }; + return Err(self.get_error()); + } + + let mut response_headers = HashMap::new(); + unsafe { + let response_header_keys: &[*mut c_char] = + slice::from_raw_parts((*response).headerKeys, (*response).headerCount as usize); + let response_header_values: &[*mut c_char] = + slice::from_raw_parts((*response).headerValues, (*response).headerCount as usize); + + for (key, value) in response_header_keys + .iter() + .zip(response_header_values.iter()) + { + response_headers.insert( + CStr::from_ptr(*key).to_str().unwrap().to_owned(), + CStr::from_ptr(*value).to_str().unwrap().to_owned(), + ); + } + } + + let r = DownloadResponse { + status_code: unsafe { (*response).statusCode }, + headers: response_headers, + }; + + unsafe { BNFreeDownloadInstanceResponse(response) }; + + Ok(r) + } +} + +// TODO: Verify the object is thread safe in the core (hint its not). +unsafe impl Send for DownloadInstance {} +unsafe impl Sync for DownloadInstance {} + +impl ToOwned for DownloadInstance { + type Owned = Ref<Self>; + + fn to_owned(&self) -> Self::Owned { + unsafe { RefCountable::inc_ref(self) } + } +} + +unsafe impl RefCountable for DownloadInstance { + unsafe fn inc_ref(handle: &Self) -> Ref<Self> { + Ref::new(Self { + handle: BNNewDownloadInstanceReference(handle.handle), + }) + } + + unsafe fn dec_ref(handle: &Self) { + BNFreeDownloadInstance(handle.handle); + } +} + +unsafe extern "C" fn cb_read_input(data: *mut u8, len: u64, ctxt: *mut c_void) -> i64 { + let callbacks = ctxt as *mut DownloadInstanceInputOutputCallbacks; + if let Some(func) = &mut (*callbacks).read { + let slice = slice::from_raw_parts_mut(data, len as usize); + let result = (func)(slice); + if let Some(count) = result { + count as i64 + } else { + -1 + } + } else { + 0 + } +} + +unsafe extern "C" fn cb_write_input(data: *mut u8, len: u64, ctxt: *mut c_void) -> u64 { + let callbacks = ctxt as *mut DownloadInstanceInputOutputCallbacks; + if let Some(func) = &mut (*callbacks).write { + let slice = slice::from_raw_parts(data, len as usize); + let result = (func)(slice); + result as u64 + } else { + 0 + } +} + +unsafe extern "C" fn cb_progress_input(ctxt: *mut c_void, progress: usize, total: usize) -> bool { + let callbacks = ctxt as *mut DownloadInstanceInputOutputCallbacks; + if let Some(func) = &mut (*callbacks).progress { + (func)(progress, total) + } else { + true + } +} + +unsafe extern "C" fn cb_write_output(data: *mut u8, len: u64, ctxt: *mut c_void) -> u64 { + let callbacks = ctxt as *mut DownloadInstanceOutputCallbacks; + if let Some(func) = &mut (*callbacks).write { + let slice = slice::from_raw_parts(data, len as usize); + let result = (func)(slice); + result as u64 + } else { + 0u64 + } +} + +unsafe extern "C" fn cb_progress_output(ctxt: *mut c_void, progress: usize, total: usize) -> bool { + let callbacks = ctxt as *mut DownloadInstanceOutputCallbacks; + if let Some(func) = &mut (*callbacks).progress { + (func)(progress, total) + } else { + true + } +} + +pub unsafe extern "C" fn cb_destroy_instance<C: CustomDownloadInstance>(ctxt: *mut c_void) { + let _ = Box::from_raw(ctxt as *mut C); +} + +pub unsafe extern "C" fn cb_perform_request<C: CustomDownloadInstance>( + ctxt: *mut c_void, + url: *const c_char, +) -> i32 { + let c = ManuallyDrop::new(Box::from_raw(ctxt as *mut C)); + + let url = match CStr::from_ptr(url).to_str() { + Ok(url) => url, + Err(e) => { + c.handle().set_error(&format!("Invalid URL: {}", e)); + return -1; + } + }; + + match c.perform_request(url) { + Ok(()) => 0, + Err(e) => { + c.handle().set_error(&e); + -1 + } + } +} + +pub unsafe extern "C" fn cb_perform_custom_request<C: CustomDownloadInstance>( + ctxt: *mut c_void, + method: *const c_char, + url: *const c_char, + header_count: u64, + header_keys: *const *const c_char, + header_values: *const *const c_char, + response: *mut *mut BNDownloadInstanceResponse, +) -> i32 { + let c = ManuallyDrop::new(Box::from_raw(ctxt as *mut C)); + + let method = match CStr::from_ptr(method).to_str() { + Ok(method) => method, + Err(e) => { + c.handle().set_error(&format!("Invalid Method: {}", e)); + return -1; + } + }; + + let url = match CStr::from_ptr(url).to_str() { + Ok(url) => url, + Err(e) => { + c.handle().set_error(&format!("Invalid URL: {}", e)); + return -1; + } + }; + + // SAFETY BnString and *mut c_char are transparent + let header_count = usize::try_from(header_count).unwrap(); + let header_keys = slice::from_raw_parts(header_keys as *const BnString, header_count); + let header_values = slice::from_raw_parts(header_values as *const BnString, header_count); + let header_keys_str = header_keys.iter().map(|s| s.to_string_lossy().to_string()); + let header_values_str = header_values + .iter() + .map(|s| s.to_string_lossy().to_string()); + let headers = header_keys_str.zip(header_values_str); + + match c.perform_custom_request(method, url, headers) { + Ok(res) => { + let res_header_keys_ptr = strings_to_string_list(res.headers.keys()); + let res_header_values_ptr = strings_to_string_list(res.headers.values()); + let raw_response = BNDownloadInstanceResponse { + statusCode: res.status_code, + headerCount: res.headers.len() as u64, + headerKeys: res_header_keys_ptr, + headerValues: res_header_values_ptr, + }; + // Leak the response and free it with cb_free_response + unsafe { *response = Box::leak(Box::new(raw_response)) }; + 0 + } + Err(e) => { + c.handle().set_error(&e); + -1 + } + } +} + +unsafe extern "C" fn cb_free_response<C: CustomDownloadInstance>( + _ctxt: *mut c_void, + response: *mut BNDownloadInstanceResponse, +) { + let _ = Box::from_raw(response); +} |
