diff options
| author | Mason Reed <mason@vector35.com> | 2026-02-19 20:25:30 -0800 |
|---|---|---|
| committer | Mason Reed <35282038+emesare@users.noreply.github.com> | 2026-02-23 00:09:44 -0800 |
| commit | dbd54d67a6d523f64615f653d82d8224cd09870a (patch) | |
| tree | 7aecf160c624e415a47fd6dc6e9a2921963c9e23 /rust/src/websocket | |
| parent | bfbb878c05220e1697aa550a23ad59911e546c17 (diff) | |
[Rust] Ensure proper lifetime management of `WebsocketClientCallback` objects
Diffstat (limited to 'rust/src/websocket')
| -rw-r--r-- | rust/src/websocket/client.rs | 162 |
1 files changed, 109 insertions, 53 deletions
diff --git a/rust/src/websocket/client.rs b/rust/src/websocket/client.rs index fc56dc0b..8095fc93 100644 --- a/rust/src/websocket/client.rs +++ b/rust/src/websocket/client.rs @@ -2,16 +2,31 @@ use crate::rc::{Ref, RefCountable}; use crate::string::{BnString, IntoCStr}; use binaryninjacore_sys::*; use std::ffi::{c_char, c_void, CStr}; +use std::marker::PhantomData; +use std::ops::Deref; use std::ptr::NonNull; pub trait WebsocketClientCallback: Sync + Send { - fn connected(&mut self) -> bool; + /// Receive a notification that the websocket connection has been connected successfully. + /// + /// Return `false` if you would like to terminate the connection early. + fn connected(&self) -> bool; - fn disconnected(&mut self); + /// Receive a notification that the websocket connection has been terminated. + /// + /// For implementations, you must call this at the end of the websocket connection lifecycle + /// even if you notify the client of an error. + fn disconnected(&self); - fn error(&mut self, msg: &str); + /// Receive an error from the websocket connection. + /// + /// For implementations, you typically write the data to an interior mutable buffer on the instance. + fn error(&self, msg: &str); - fn read(&mut self, data: &[u8]) -> bool; + /// Receive data from the websocket connection. + /// + /// For implementations, you typically write the data to an interior mutable buffer on the instance. + fn read(&self, data: &[u8]) -> bool; } pub trait WebsocketClient: Sync + Send { @@ -27,7 +42,32 @@ pub trait WebsocketClient: Sync + Send { fn disconnect(&self) -> bool; } +/// Represents a live websocket connection. +/// +/// This manages the lifetime of the callback, ensuring it outlives the connection. +pub struct ActiveConnection<'a, C: WebsocketClientCallback> { + pub client: Ref<CoreWebsocketClient>, + _callback: PhantomData<&'a mut C>, +} + +impl<'a, C: WebsocketClientCallback> Deref for ActiveConnection<'a, C> { + type Target = CoreWebsocketClient; + + fn deref(&self) -> &Self::Target { + &self.client + } +} + +impl<'a, C: WebsocketClientCallback> Drop for ActiveConnection<'a, C> { + fn drop(&mut self) { + self.client.disconnect(); + } +} + /// Implements a websocket client. +/// +/// To connect, use [`Ref<CoreWebsocketClient>::connect`] which will return an [`ActiveConnection`] +/// which manages the lifecycle of the websocket connection. #[repr(transparent)] pub struct CoreWebsocketClient { pub(crate) handle: NonNull<BNWebsocketClient>, @@ -43,7 +83,50 @@ impl CoreWebsocketClient { &mut *self.handle.as_ptr() } - /// Initializes the web socket connection. + /// 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. This must be called at the end of an active + /// websocket connection lifecycle, to free resources. + /// + /// 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, this is not a terminating request you must use + /// [`CoreWebsocketClient::notify_disconnected`] to terminate the connection. + pub fn notify_error(&self, msg: &str) { + let error = msg.to_cstr(); + unsafe { BNNotifyWebsocketClientError(self.handle.as_ptr(), error.as_ptr()) } + } + + /// 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()) } + } +} + +impl Ref<CoreWebsocketClient> { + /// Initializes the web socket connection, returning the [`ActiveConnection`], once dropped the + /// connection will be disconnected. /// /// Connect to a given url, asynchronously. The connection will be run in a /// separate thread managed by the websocket provider. @@ -54,20 +137,22 @@ impl CoreWebsocketClient { /// 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 succeeds but later fails, [`WebsocketClientCallback::error`] will be called + /// and shortly thereafter [`WebsocketClientCallback::disconnected`] will be called. /// - /// 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 + /// 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, C>(&self, host: &str, headers: I, callbacks: &mut C) -> bool + #[must_use] + pub fn connect<'a, I, C>( + self, + host: &str, + headers: I, + callbacks: &'a C, + ) -> Option<ActiveConnection<'a, C>> where I: IntoIterator<Item = (String, String)>, C: WebsocketClientCallback, @@ -79,16 +164,16 @@ impl CoreWebsocketClient { .unzip(); let header_keys: Vec<*const c_char> = header_keys.iter().map(|k| k.as_ptr()).collect(); let header_values: Vec<*const c_char> = header_values.iter().map(|v| v.as_ptr()).collect(); - // SAFETY: This context will only be live for the duration of BNConnectWebsocketClient + // SAFETY: This context will live for as long as the `ActiveConnection` is alive. // 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, + context: callbacks as *const C 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 { + let success = unsafe { BNConnectWebsocketClient( self.handle.as_ptr(), url.as_ptr(), @@ -97,46 +182,17 @@ impl CoreWebsocketClient { 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.to_cstr(); - unsafe { BNNotifyWebsocketClientError(self.handle.as_ptr(), error.as_ptr()) } - } + }; - /// 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(), - ) + if success { + Some(ActiveConnection { + client: self, + _callback: PhantomData, + }) + } else { + None } } - - 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 {} |
