summaryrefslogtreecommitdiff
path: root/rust
diff options
context:
space:
mode:
authorRubens Brandao <git@rubens.io>2024-06-26 17:24:33 -0300
committerMason Reed <mason@vector35.com>2025-02-07 15:52:59 -0500
commita039d392b54170d236d9d993024e9eca38f701ce (patch)
treeac5900027386c9b4a1197689008c646eb59e0b80 /rust
parent3f0a217d7d097c0502af81e78a01ed3bf9f38761 (diff)
Implement Rust WebsocketProvider
Diffstat (limited to 'rust')
-rw-r--r--rust/src/lib.rs1
-rw-r--r--rust/src/websocket.rs10
-rw-r--r--rust/src/websocket/client.rs245
-rw-r--r--rust/src/websocket/provider.rs128
-rw-r--r--rust/tests/websocket.rs124
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);
+}