diff options
| -rw-r--r-- | rust/src/lib.rs | 1 | ||||
| -rw-r--r-- | rust/src/websocket.rs | 10 | ||||
| -rw-r--r-- | rust/src/websocket/client.rs | 245 | ||||
| -rw-r--r-- | rust/src/websocket/provider.rs | 128 | ||||
| -rw-r--r-- | rust/tests/websocket.rs | 124 |
5 files changed, 508 insertions, 0 deletions
diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 34c6bb53..d8d47938 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -86,6 +86,7 @@ pub mod type_printer; pub mod types; pub mod update; pub mod variable; +pub mod websocket; pub mod worker_thread; pub mod workflow; diff --git a/rust/src/websocket.rs b/rust/src/websocket.rs new file mode 100644 index 00000000..cc1b5f77 --- /dev/null +++ b/rust/src/websocket.rs @@ -0,0 +1,10 @@ +//! Interface for registering new websocket providers +//! +//! WARNING: Do _not_ use this for anything other than provider registration. If you need to open a +//! websocket connection use a real websocket library. + +mod client; +mod provider; + +pub use client::*; +pub use provider::*; diff --git a/rust/src/websocket/client.rs b/rust/src/websocket/client.rs new file mode 100644 index 00000000..36c7bedd --- /dev/null +++ b/rust/src/websocket/client.rs @@ -0,0 +1,245 @@ +use crate::rc::{Ref, RefCountable}; +use crate::string::{BnStrCompatible, BnString}; +use binaryninjacore_sys::*; +use std::ffi::{c_char, c_void, CStr}; +use std::ptr::NonNull; + +pub trait WebsocketClientCallback: Sync + Send { + fn connected(&mut self) -> bool; + + fn disconnected(&mut self); + + fn error(&mut self, msg: &str); + + fn read(&mut self, data: &[u8]) -> bool; +} + +pub trait WebsocketClient: Sync + Send { + /// Called to construct this client object with the given core object. + fn from_core(core: Ref<CoreWebsocketClient>) -> Self; + + fn connect<I, K, V>(&self, host: &str, headers: I) -> bool + where + I: IntoIterator<Item = (K, V)>, + K: BnStrCompatible, + V: BnStrCompatible; + + fn write(&self, data: &[u8]) -> bool; + + fn disconnect(&self) -> bool; +} + +/// Implements a websocket client. +#[repr(transparent)] +pub struct CoreWebsocketClient { + pub(crate) handle: NonNull<BNWebsocketClient>, +} + +impl CoreWebsocketClient { + pub(crate) unsafe fn ref_from_raw(handle: NonNull<BNWebsocketClient>) -> Ref<Self> { + Ref::new(Self { handle }) + } + + #[allow(clippy::mut_from_ref)] + pub(crate) unsafe fn as_raw(&self) -> &mut BNWebsocketClient { + &mut *self.handle.as_ptr() + } + + /// Initializes the web socket connection. + /// + /// Connect to a given url, asynchronously. The connection will be run in a + /// separate thread managed by the websocket provider. + /// + /// Callbacks will be called **on the thread of the connection**, so be sure + /// to ExecuteOnMainThread any long-running or gui operations in the callbacks. + /// + /// If the connection succeeds, [WebsocketClientCallback::connected] will be called. On normal + /// termination, [WebsocketClientCallback::disconnected] will be called. + /// + /// If the connection succeeds, but later fails, [WebsocketClientCallback::disconnected] will not + /// be called, and [WebsocketClientCallback::error] will be called instead. + /// + /// If the connection fails, neither [WebsocketClientCallback::connected] nor + /// [WebsocketClientCallback::disconnected] will be called, and [WebsocketClientCallback::error] + /// will be called instead. + /// + /// If [WebsocketClientCallback::connected] or [WebsocketClientCallback::read] return false, the + /// connection will be aborted. + /// + /// * `host` - Full url with scheme, domain, optionally port, and path + /// * `headers` - HTTP header keys and values + /// * `callback` - Callbacks for various websocket events + pub fn initialize_connection<I, K, V, C>( + &self, + host: &str, + headers: I, + callbacks: &mut C, + ) -> bool + where + I: IntoIterator<Item = (K, V)>, + K: BnStrCompatible, + V: BnStrCompatible, + C: WebsocketClientCallback, + { + let url = host.into_bytes_with_nul(); + let (header_keys, header_values): (Vec<K::Result>, Vec<V::Result>) = headers + .into_iter() + .map(|(k, v)| (k.into_bytes_with_nul(), v.into_bytes_with_nul())) + .unzip(); + let header_keys: Vec<*const c_char> = header_keys + .iter() + .map(|k| k.as_ref().as_ptr() as *const c_char) + .collect(); + let header_values: Vec<*const c_char> = header_values + .iter() + .map(|v| v.as_ref().as_ptr() as *const c_char) + .collect(); + // SAFETY: This context will only be live for the duration of BNConnectWebsocketClient + // SAFETY: Any subsequent call to BNConnectWebsocketClient will write over the context. + let mut output_callbacks = BNWebsocketClientOutputCallbacks { + context: callbacks as *mut C as *mut c_void, + connectedCallback: Some(cb_connected::<C>), + disconnectedCallback: Some(cb_disconnected::<C>), + errorCallback: Some(cb_error::<C>), + readCallback: Some(cb_read::<C>), + }; + unsafe { + BNConnectWebsocketClient( + self.handle.as_ptr(), + url.as_ptr() as *const c_char, + header_keys.len().try_into().unwrap(), + header_keys.as_ptr(), + header_values.as_ptr(), + &mut output_callbacks, + ) + } + } + + /// Call the connect callback function, forward the callback returned value + pub fn notify_connected(&self) -> bool { + unsafe { BNNotifyWebsocketClientConnect(self.handle.as_ptr()) } + } + + /// Notify the callback function of a disconnect, + /// + /// NOTE: This does not actually disconnect, use the [Self::disconnect] function for that. + pub fn notify_disconnected(&self) { + unsafe { BNNotifyWebsocketClientDisconnect(self.handle.as_ptr()) } + } + + /// Call the error callback function + pub fn notify_error(&self, msg: &str) { + let error = msg.into_bytes_with_nul(); + unsafe { + BNNotifyWebsocketClientError(self.handle.as_ptr(), error.as_ptr() as *const c_char) + } + } + + /// Call the read callback function, forward the callback returned value + pub fn notify_read(&self, data: &[u8]) -> bool { + unsafe { + BNNotifyWebsocketClientReadData( + self.handle.as_ptr(), + data.as_ptr() as *mut _, + data.len().try_into().unwrap(), + ) + } + } + + pub fn write(&self, data: &[u8]) -> bool { + let len = u64::try_from(data.len()).unwrap(); + unsafe { BNWriteWebsocketClientData(self.as_raw(), data.as_ptr(), len) != 0 } + } + + pub fn disconnect(&self) -> bool { + unsafe { BNDisconnectWebsocketClient(self.as_raw()) } + } +} + +unsafe impl Sync for CoreWebsocketClient {} +unsafe impl Send for CoreWebsocketClient {} + +impl ToOwned for CoreWebsocketClient { + type Owned = Ref<Self>; + + fn to_owned(&self) -> Self::Owned { + unsafe { RefCountable::inc_ref(self) } + } +} + +unsafe impl RefCountable for CoreWebsocketClient { + unsafe fn inc_ref(handle: &Self) -> Ref<Self> { + let result = BNNewWebsocketClientReference(handle.as_raw()); + unsafe { Self::ref_from_raw(NonNull::new(result).unwrap()) } + } + + unsafe fn dec_ref(handle: &Self) { + BNFreeWebsocketClient(handle.as_raw()) + } +} + +pub(crate) unsafe extern "C" fn cb_destroy_client<W: WebsocketClient>(ctxt: *mut c_void) { + let _ = Box::from_raw(ctxt as *mut W); +} + +pub(crate) unsafe extern "C" fn cb_connect<W: WebsocketClient>( + ctxt: *mut c_void, + host: *const c_char, + header_count: u64, + header_keys: *const *const c_char, + header_values: *const *const c_char, +) -> bool { + let ctxt: &mut W = &mut *(ctxt as *mut W); + let host = CStr::from_ptr(host); + // SAFETY BnString and *mut c_char are transparent + let header_count = usize::try_from(header_count).unwrap(); + let header_keys = core::slice::from_raw_parts(header_keys as *const BnString, header_count); + let header_values = core::slice::from_raw_parts(header_values as *const BnString, header_count); + let header_keys_str = header_keys.iter().map(|s| s.to_string_lossy()); + let header_values_str = header_values.iter().map(|s| s.to_string_lossy()); + let header = header_keys_str.zip(header_values_str); + ctxt.connect(&host.to_string_lossy(), header) +} + +pub(crate) unsafe extern "C" fn cb_write<W: WebsocketClient>( + data: *const u8, + len: u64, + ctxt: *mut c_void, +) -> bool { + let ctxt: &mut W = &mut *(ctxt as *mut W); + let len = usize::try_from(len).unwrap(); + let data = core::slice::from_raw_parts(data, len); + ctxt.write(data) +} + +pub(crate) unsafe extern "C" fn cb_disconnect<W: WebsocketClient>(ctxt: *mut c_void) -> bool { + let ctxt: &mut W = &mut *(ctxt as *mut W); + ctxt.disconnect() +} + +unsafe extern "C" fn cb_connected<W: WebsocketClientCallback>(ctxt: *mut c_void) -> bool { + let ctxt: &mut W = &mut *(ctxt as *mut W); + ctxt.connected() +} + +unsafe extern "C" fn cb_disconnected<W: WebsocketClientCallback>(ctxt: *mut c_void) { + let ctxt: &mut W = &mut *(ctxt as *mut W); + ctxt.disconnected() +} + +unsafe extern "C" fn cb_error<W: WebsocketClientCallback>(msg: *const c_char, ctxt: *mut c_void) { + let ctxt: &mut W = &mut *(ctxt as *mut W); + let msg = CStr::from_ptr(msg); + ctxt.error(&msg.to_string_lossy()) +} + +unsafe extern "C" fn cb_read<W: WebsocketClientCallback>( + data: *mut u8, + len: u64, + ctxt: *mut c_void, +) -> bool { + let ctxt: &mut W = &mut *(ctxt as *mut W); + let len = usize::try_from(len).unwrap(); + let data = core::slice::from_raw_parts_mut(data, len); + ctxt.read(data) +} diff --git a/rust/src/websocket/provider.rs b/rust/src/websocket/provider.rs new file mode 100644 index 00000000..0e28afe4 --- /dev/null +++ b/rust/src/websocket/provider.rs @@ -0,0 +1,128 @@ +use crate::rc::{Array, CoreArrayProvider, CoreArrayProviderInner, Ref}; +use crate::string::{BnStrCompatible, BnString}; +use crate::websocket::client; +use crate::websocket::client::{CoreWebsocketClient, WebsocketClient}; +use binaryninjacore_sys::*; +use std::ffi::{c_char, c_void}; +use std::mem::MaybeUninit; +use std::ptr::NonNull; + +pub fn register_websocket_provider<W>(name: &str) -> &'static mut W +where + W: WebsocketProvider, +{ + let name = name.into_bytes_with_nul(); + let provider_uninit = MaybeUninit::uninit(); + // SAFETY: Websocket provider is never freed + let leaked_provider = Box::leak(Box::new(provider_uninit)); + let result = unsafe { + BNRegisterWebsocketProvider( + name.as_ptr() as *const c_char, + &mut BNWebsocketProviderCallbacks { + context: leaked_provider as *mut _ as *mut c_void, + createClient: Some(cb_create_client::<W>), + }, + ) + }; + + let provider_core = unsafe { CoreWebsocketProvider::from_raw(NonNull::new(result).unwrap()) }; + // We now have the core provider so we can actually construct the object. + leaked_provider.write(W::from_core(provider_core)); + unsafe { leaked_provider.assume_init_mut() } +} + +pub trait WebsocketProvider: Sync + Send + Sized { + type Client: WebsocketClient; + + fn handle(&self) -> CoreWebsocketProvider; + + /// Called to construct this provider object with the given core object. + fn from_core(core: CoreWebsocketProvider) -> Self; + + /// Create a new instance of the websocket client. + fn create_client(&self) -> Result<Ref<CoreWebsocketClient>, ()> { + let client_uninit = MaybeUninit::uninit(); + // SAFETY: Websocket client is freed by cb_destroy_client + let leaked_client = Box::leak(Box::new(client_uninit)); + let mut callbacks = BNWebsocketClientCallbacks { + context: leaked_client as *mut _ as *mut c_void, + connect: Some(client::cb_connect::<Self::Client>), + destroyClient: Some(client::cb_destroy_client::<Self::Client>), + disconnect: Some(client::cb_disconnect::<Self::Client>), + write: Some(client::cb_write::<Self::Client>), + }; + let client_ptr = + unsafe { BNInitWebsocketClient(self.handle().handle.as_ptr(), &mut callbacks) }; + // TODO: If possible pass a sensible error back... + let client_ptr = NonNull::new(client_ptr).ok_or(())?; + let client_ref = unsafe { CoreWebsocketClient::ref_from_raw(client_ptr) }; + // We now have the core client so we can actually construct the object. + leaked_client.write(Self::Client::from_core(client_ref.clone())); + Ok(client_ref) + } +} + +#[derive(Clone, Copy, Hash, PartialEq, Eq)] +#[repr(transparent)] +pub struct CoreWebsocketProvider { + handle: NonNull<BNWebsocketProvider>, +} + +impl CoreWebsocketProvider { + pub(crate) unsafe fn from_raw(handle: NonNull<BNWebsocketProvider>) -> Self { + Self { handle } + } + + pub fn all() -> Array<Self> { + let mut count = 0; + let result = unsafe { BNGetWebsocketProviderList(&mut count) }; + assert!(!result.is_null()); + unsafe { Array::new(result, count, ()) } + } + + pub fn by_name<S: BnStrCompatible>(name: S) -> Option<CoreWebsocketProvider> { + let name = name.into_bytes_with_nul(); + let result = + unsafe { BNGetWebsocketProviderByName(name.as_ref().as_ptr() as *const c_char) }; + NonNull::new(result).map(|h| unsafe { Self::from_raw(h) }) + } + + pub fn name(&self) -> BnString { + let result = unsafe { BNGetWebsocketProviderName(self.handle.as_ptr()) }; + assert!(!result.is_null()); + unsafe { BnString::from_raw(result) } + } +} + +unsafe impl Sync for CoreWebsocketProvider {} +unsafe impl Send for CoreWebsocketProvider {} + +impl CoreArrayProvider for CoreWebsocketProvider { + type Raw = *mut BNWebsocketProvider; + type Context = (); + type Wrapped<'a> = Self; +} + +unsafe impl CoreArrayProviderInner for CoreWebsocketProvider { + unsafe fn free(raw: *mut Self::Raw, _count: usize, _context: &Self::Context) { + BNFreeWebsocketProviderList(raw) + } + + unsafe fn wrap_raw<'a>(raw: &'a Self::Raw, _context: &'a Self::Context) -> Self::Wrapped<'a> { + let handle = NonNull::new(*raw).unwrap(); + Self::from_raw(handle) + } +} + +unsafe extern "C" fn cb_create_client<W: WebsocketProvider>( + ctxt: *mut c_void, +) -> *mut BNWebsocketClient { + let ctxt: &mut W = &mut *(ctxt as *mut W); + match ctxt.create_client() { + Ok(owned_client) => { + // SAFETY: The caller is assumed to have picked up this ref. + Ref::into_raw(owned_client).handle.as_ptr() + } + Err(_) => std::ptr::null_mut(), + } +} diff --git a/rust/tests/websocket.rs b/rust/tests/websocket.rs new file mode 100644 index 00000000..9feb9381 --- /dev/null +++ b/rust/tests/websocket.rs @@ -0,0 +1,124 @@ +use binaryninja::headless::Session; +use binaryninja::rc::Ref; +use binaryninja::string::BnStrCompatible; +use binaryninja::websocket::{ + register_websocket_provider, CoreWebsocketClient, CoreWebsocketProvider, WebsocketClient, + WebsocketClientCallback, WebsocketProvider, +}; +use rstest::*; + +#[fixture] +#[once] +fn session() -> Session { + Session::new().expect("Failed to initialize session") +} + +struct MyWebsocketProvider { + core: CoreWebsocketProvider, +} + +impl WebsocketProvider for MyWebsocketProvider { + type Client = MyWebsocketClient; + + fn handle(&self) -> CoreWebsocketProvider { + self.core + } + + fn from_core(core: CoreWebsocketProvider) -> Self { + MyWebsocketProvider { core } + } +} + +struct MyWebsocketClient { + core: Ref<CoreWebsocketClient>, +} + +impl WebsocketClient for MyWebsocketClient { + fn from_core(core: Ref<CoreWebsocketClient>) -> Self { + Self { core } + } + + fn connect<I, K, V>(&self, host: &str, _headers: I) -> bool + where + I: IntoIterator<Item = (K, V)>, + K: BnStrCompatible, + V: BnStrCompatible, + { + assert_eq!(host, "url"); + true + } + + fn write(&self, data: &[u8]) -> bool { + if !self.core.notify_read("sent: ".as_bytes()) { + return false; + } + if !self.core.notify_read(data) { + return false; + } + self.core.notify_read("\n".as_bytes()) + } + + fn disconnect(&self) -> bool { + true + } +} + +#[derive(Default)] +struct MyClientCallbacks { + data_read: Vec<u8>, + did_disconnect: bool, + did_error: bool, +} + +impl WebsocketClientCallback for MyClientCallbacks { + fn connected(&mut self) -> bool { + true + } + + fn disconnected(&mut self) { + self.did_disconnect = true; + } + + fn error(&mut self, msg: &str) { + assert_eq!(msg, "error"); + self.did_error = true; + } + + fn read(&mut self, data: &[u8]) -> bool { + self.data_read.extend_from_slice(data); + true + } +} + +#[rstest] +fn reg_websocket_provider(_session: &Session) { + let provider = register_websocket_provider::<MyWebsocketProvider>("RustWebsocketProvider"); + let client = provider.create_client().unwrap(); + let mut callback = MyClientCallbacks::default(); + let success = client.initialize_connection("url", [("header", "value")], &mut callback); + assert!(success, "Failed to initialize connection!"); +} + +#[rstest] +fn listen_websocket_provider(_session: &Session) { + let provider = register_websocket_provider::<MyWebsocketProvider>("RustWebsocketProvider2"); + + let client = provider.create_client().unwrap(); + let mut callback = MyClientCallbacks::default(); + client.initialize_connection("url", [("header", "value")], &mut callback); + + assert!(client.write("test1".as_bytes())); + assert!(client.write("test2".as_bytes())); + + client.notify_error("error"); + client.disconnect(); + drop(client); + + assert_eq!( + &callback.data_read[..], + "sent: test1\nsent: test2\n".as_bytes() + ); + // If we disconnected that means the error callback was not notified. + assert!(!callback.did_disconnect); + assert!(callback.did_error); +} |
