diff --git a/crates/livekit-webrtc/libwebrtc-sys/include/livekit/data_channel.h b/crates/livekit-webrtc/libwebrtc-sys/include/livekit/data_channel.h index e1b2211..52c9487 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/include/livekit/data_channel.h +++ b/crates/livekit-webrtc/libwebrtc-sys/include/livekit/data_channel.h @@ -20,6 +20,8 @@ namespace livekit { void register_observer(NativeDataChannelObserver &observer); void unregister_observer(); + bool send(const DataBuffer& buffer); + rust::String label() const; void close(); private: rtc::scoped_refptr data_channel_; diff --git a/crates/livekit-webrtc/libwebrtc-sys/include/livekit/rust_types.h b/crates/livekit-webrtc/libwebrtc-sys/include/livekit/rust_types.h index d9334bf..9e42d9d 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/include/livekit/rust_types.h +++ b/crates/livekit-webrtc/libwebrtc-sys/include/livekit/rust_types.h @@ -21,6 +21,7 @@ namespace livekit { struct RTCOfferAnswerOptions; struct RTCError; struct DataChannelInit; + struct DataBuffer; } #endif //RUST_TYPES_H diff --git a/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.cpp b/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.cpp index 1da192e..ad031ec 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.cpp +++ b/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.cpp @@ -21,6 +21,14 @@ namespace livekit { data_channel_->UnregisterObserver(); } + bool DataChannel::send(const DataBuffer &buffer) { + return data_channel_->Send(webrtc::DataBuffer{rtc::CopyOnWriteBuffer(buffer.ptr, buffer.len), buffer.binary }); + } + + rust::String DataChannel::label() const{ + return data_channel_->label(); + } + void DataChannel::close() { return data_channel_->Close(); } @@ -55,7 +63,7 @@ namespace livekit { void NativeDataChannelObserver::OnMessage(const webrtc::DataBuffer &buffer) { DataBuffer data{}; - data.binary = buffer.data.data(); + data.ptr = buffer.data.data(); data.len = buffer.data.size(); data.binary = buffer.binary; observer_->on_message(data); diff --git a/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.rs b/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.rs index 5b0fa35..fc1854b 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.rs +++ b/crates/livekit-webrtc/libwebrtc-sys/src/data_channel.rs @@ -38,6 +38,7 @@ pub mod ffi { } #[derive(Debug)] + #[repr(u32)] pub enum DataState { Connecting, Open, @@ -68,6 +69,8 @@ pub mod ffi { ); fn unregister_observer(self: Pin<&mut DataChannel>); + fn send(self: Pin<&mut DataChannel>, data: &DataBuffer) -> bool; + fn label(self: &DataChannel) -> String; fn close(self: Pin<&mut DataChannel>); fn create_data_channel_init(init: DataChannelInit) -> UniquePtr; @@ -79,6 +82,9 @@ pub mod ffi { } } +unsafe impl Send for ffi::DataChannel {} +unsafe impl Send for ffi::NativeDataChannelObserver {} + // DataChannelObserver pub trait DataChannelObserver: Send { @@ -88,24 +94,30 @@ pub trait DataChannelObserver: Send { } pub struct DataChannelObserverWrapper { - observer: Box, + observer: *mut dyn DataChannelObserver, } impl DataChannelObserverWrapper { - pub fn new(observer: Box) -> Self { + /// SAFETY + /// DataChannelObserver must lives as long as DataChannelObserverWrapper does + pub unsafe fn new(observer: *mut dyn DataChannelObserver) -> Self { Self { observer } } fn on_state_change(&self) { - self.observer.on_state_change(); + unsafe { + (*self.observer).on_state_change(); + } } fn on_message(&self, buffer: ffi::DataBuffer) { - let data = unsafe { slice::from_raw_parts(buffer.ptr, buffer.len) }; - self.observer.on_message(data, buffer.binary); + unsafe { + let data = slice::from_raw_parts(buffer.ptr, buffer.len); + (*self.observer).on_message(data, buffer.binary); + } } fn on_buffered_amount_change(&self, sent_data_size: u64) { - self.observer.on_buffered_amount_change(sent_data_size); + unsafe { (*self.observer).on_buffered_amount_change(sent_data_size) }; } } diff --git a/crates/livekit-webrtc/libwebrtc-sys/src/media_stream_interface.cpp b/crates/livekit-webrtc/libwebrtc-sys/src/media_stream_interface.cpp index 94487ea..c8c9a5a 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/src/media_stream_interface.cpp +++ b/crates/livekit-webrtc/libwebrtc-sys/src/media_stream_interface.cpp @@ -6,7 +6,7 @@ namespace livekit { - MediaStreamInterface::MediaStreamInterface(rtc::scoped_refptr stream) : media_stream_(stream) { + MediaStreamInterface::MediaStreamInterface(rtc::scoped_refptr stream) : media_stream_(std::move(stream)) { } } // livekit \ No newline at end of file diff --git a/crates/livekit-webrtc/libwebrtc-sys/src/peer_connection.rs b/crates/livekit-webrtc/libwebrtc-sys/src/peer_connection.rs index ae6516a..72f6424 100644 --- a/crates/livekit-webrtc/libwebrtc-sys/src/peer_connection.rs +++ b/crates/livekit-webrtc/libwebrtc-sys/src/peer_connection.rs @@ -2,10 +2,10 @@ use crate::candidate::ffi::Candidate; use crate::data_channel::ffi::DataChannel; use crate::jsep::ffi::IceCandidate; use crate::media_stream_interface::ffi::MediaStreamInterface; +use crate::rtc_error::ffi::RTCError; use crate::rtp_receiver::ffi::RtpReceiver; use crate::rtp_transceiver::ffi::RtpTransceiver; use cxx::UniquePtr; -use crate::rtc_error::ffi::RTCError; #[cxx::bridge(namespace = "livekit")] pub mod ffi { @@ -174,17 +174,11 @@ pub mod ffi { extern "Rust" { type AddIceCandidateObserverWrapper; - fn on_complete( - self: &AddIceCandidateObserverWrapper, - error: RTCError, - ); + fn on_complete(self: &AddIceCandidateObserverWrapper, error: RTCError); type PeerConnectionObserverWrapper; - fn on_signaling_change( - self: &PeerConnectionObserverWrapper, - new_state: SignalingState, - ); + fn on_signaling_change(self: &PeerConnectionObserverWrapper, new_state: SignalingState); fn on_add_stream( self: &PeerConnectionObserverWrapper, stream: UniquePtr, @@ -244,14 +238,8 @@ pub mod ffi { receiver: UniquePtr, streams: Vec, ); - fn on_track( - self: &PeerConnectionObserverWrapper, - transceiver: UniquePtr, - ); - fn on_remove_track( - self: &PeerConnectionObserverWrapper, - receiver: UniquePtr, - ); + fn on_track(self: &PeerConnectionObserverWrapper, transceiver: UniquePtr); + fn on_remove_track(self: &PeerConnectionObserverWrapper, receiver: UniquePtr); fn on_interesting_usage(self: &PeerConnectionObserverWrapper, usage_pattern: i32); } } @@ -260,6 +248,8 @@ pub mod ffi { unsafe impl Sync for ffi::PeerConnection {} unsafe impl Send for ffi::PeerConnection {} +unsafe impl Send for ffi::NativePeerConnectionObserver {} + impl Default for ffi::RTCOfferAnswerOptions { /* static const int kUndefined = -1; @@ -281,9 +271,8 @@ impl Default for ffi::RTCOfferAnswerOptions { } } - pub struct AddIceCandidateObserverWrapper { - observer: Box + observer: Box, } impl AddIceCandidateObserverWrapper { @@ -291,7 +280,7 @@ impl AddIceCandidateObserverWrapper { Self { observer } } - fn on_complete(&self, error: RTCError){ + fn on_complete(&self, error: RTCError) { (self.observer)(error); } } @@ -342,47 +331,69 @@ impl PeerConnectionObserverWrapper { } fn on_signaling_change(&self, new_state: ffi::SignalingState) { - unsafe { (*self.observer).on_signaling_change(new_state); } + unsafe { + (*self.observer).on_signaling_change(new_state); + } } fn on_add_stream(&self, stream: UniquePtr) { - unsafe { (*self.observer).on_add_stream(stream); } + unsafe { + (*self.observer).on_add_stream(stream); + } } fn on_remove_stream(&self, stream: UniquePtr) { - unsafe { (*self.observer).on_remove_stream(stream); } + unsafe { + (*self.observer).on_remove_stream(stream); + } } fn on_data_channel(&self, data_channel: UniquePtr) { - unsafe { (*self.observer).on_data_channel(data_channel); } + unsafe { + (*self.observer).on_data_channel(data_channel); + } } fn on_renegotiation_needed(&self) { - unsafe { (*self.observer).on_renegotiation_needed(); } + unsafe { + (*self.observer).on_renegotiation_needed(); + } } fn on_negotiation_needed_event(&self, event: u32) { - unsafe { (*self.observer).on_negotiation_needed_event(event); } + unsafe { + (*self.observer).on_negotiation_needed_event(event); + } } fn on_ice_connection_change(&self, new_state: ffi::IceConnectionState) { - unsafe { (*self.observer).on_ice_connection_change(new_state); } + unsafe { + (*self.observer).on_ice_connection_change(new_state); + } } fn on_standardized_ice_connection_change(&self, new_state: ffi::IceConnectionState) { - unsafe { (*self.observer).on_standardized_ice_connection_change(new_state); } + unsafe { + (*self.observer).on_standardized_ice_connection_change(new_state); + } } fn on_connection_change(&self, new_state: ffi::PeerConnectionState) { - unsafe { (*self.observer).on_connection_change(new_state); } + unsafe { + (*self.observer).on_connection_change(new_state); + } } fn on_ice_gathering_change(&self, new_state: ffi::IceGatheringState) { - unsafe { (*self.observer).on_ice_gathering_change(new_state); } + unsafe { + (*self.observer).on_ice_gathering_change(new_state); + } } fn on_ice_candidate(&self, candidate: UniquePtr) { - unsafe { (*self.observer).on_ice_candidate(candidate); } + unsafe { + (*self.observer).on_ice_candidate(candidate); + } } fn on_ice_candidate_error( @@ -393,7 +404,9 @@ impl PeerConnectionObserverWrapper { error_code: i32, error_text: String, ) { - unsafe { (*self.observer).on_ice_candidate_error(address, port, url, error_code, error_text); } + unsafe { + (*self.observer).on_ice_candidate_error(address, port, url, error_code, error_text); + } } fn on_ice_candidates_removed(&self, removed: Vec) { @@ -403,40 +416,50 @@ impl PeerConnectionObserverWrapper { vec.push(v.ptr); } - unsafe { (*self.observer).on_ice_candidates_removed(vec); } + unsafe { + (*self.observer).on_ice_candidates_removed(vec); + } } fn on_ice_connection_receiving_change(&self, receiving: bool) { - unsafe { (*self.observer).on_ice_connection_receiving_change(receiving); } + unsafe { + (*self.observer).on_ice_connection_receiving_change(receiving); + } } fn on_ice_selected_candidate_pair_changed(&self, event: ffi::CandidatePairChangeEvent) { - unsafe { (*self.observer).on_ice_selected_candidate_pair_changed(event); } + unsafe { + (*self.observer).on_ice_selected_candidate_pair_changed(event); + } } - fn on_add_track( - &self, - receiver: UniquePtr, - streams: Vec, - ) { + fn on_add_track(&self, receiver: UniquePtr, streams: Vec) { let mut vec = Vec::new(); for v in streams { vec.push(v.ptr); } - unsafe { (*self.observer).on_add_track(receiver, vec); } + unsafe { + (*self.observer).on_add_track(receiver, vec); + } } fn on_track(&self, transceiver: UniquePtr) { - unsafe { (*self.observer).on_track(transceiver); } + unsafe { + (*self.observer).on_track(transceiver); + } } fn on_remove_track(&self, receiver: UniquePtr) { - unsafe { (*self.observer).on_remove_track(receiver); } + unsafe { + (*self.observer).on_remove_track(receiver); + } } fn on_interesting_usage(&self, usage_pattern: i32) { - unsafe { (*self.observer).on_interesting_usage(usage_pattern); } + unsafe { + (*self.observer).on_interesting_usage(usage_pattern); + } } } diff --git a/crates/livekit-webrtc/src/data_channel.rs b/crates/livekit-webrtc/src/data_channel.rs index a1af993..e87f6ea 100644 --- a/crates/livekit-webrtc/src/data_channel.rs +++ b/crates/livekit-webrtc/src/data_channel.rs @@ -1,15 +1,125 @@ +use std::fmt::{Debug, Formatter}; use cxx::UniquePtr; use libwebrtc_sys::data_channel as sys_dc; +use log::trace; +use std::sync::{Arc, Mutex}; pub use sys_dc::ffi::Priority; pub struct DataChannel { cxx_handle: UniquePtr, + observer: Box, + + // Keep alive for C++ + native_observer: UniquePtr, +} + +impl Debug for DataChannel { + fn fmt(&self, f: &mut Formatter) -> std::fmt::Result { + write!(f, "DataChannel [{:?}]", self.label()) + } } impl DataChannel { pub(crate) fn new(cxx_handle: UniquePtr) -> Self { - Self { cxx_handle } + let mut observer = Box::new(InternalDataChannelObserver::default()); + + let mut dc = unsafe { + Self { + cxx_handle, + native_observer: sys_dc::ffi::create_native_data_channel_observer(Box::new( + sys_dc::DataChannelObserverWrapper::new(&mut *observer), + )), + observer, + } + }; + + unsafe { + dc.cxx_handle + .pin_mut() + .register_observer(dc.native_observer.pin_mut()); + } + + dc + } + + pub fn send(&mut self, data: &[u8], binary: bool) -> bool { + let buffer = sys_dc::ffi::DataBuffer { + ptr: data.as_ptr(), + len: data.len(), + binary, + }; + self.cxx_handle.pin_mut().send(&buffer) + } + + pub fn label(&self) -> String { + self.cxx_handle.label() + } + + pub fn close(&mut self) { + self.cxx_handle.pin_mut().close(); + } + + pub fn on_state_change(&mut self, handler: OnStateChangeHandler) { + *self.observer.on_state_change_handler.lock().unwrap() = Some(handler); + } + + pub fn on_message(&mut self, handler: OnMessageHandler) { + *self.observer.on_message_handler.lock().unwrap() = Some(handler); + } + + pub fn on_buffer(&mut self, handler: OnBufferedAmountChangeHandler) { + *self + .observer + .on_buffered_amount_change_handler + .lock() + .unwrap() = Some(handler); + } +} + +pub type OnStateChangeHandler = Box; +pub type OnMessageHandler = Box; // data, is_binary +pub type OnBufferedAmountChangeHandler = Box; + +struct InternalDataChannelObserver { + on_state_change_handler: Arc>>, + on_message_handler: Arc>>, + on_buffered_amount_change_handler: Arc>>, +} + +impl sys_dc::DataChannelObserver for InternalDataChannelObserver { + fn on_state_change(&self) { + trace!("DataChannel: on_state_change"); + let mut handler = self.on_state_change_handler.lock().unwrap(); + if let Some(f) = handler.as_mut() { + f(); + } + } + + fn on_message(&self, data: &[u8], is_binary: bool) { + trace!("DataChannel: on_message"); + let mut handler = self.on_message_handler.lock().unwrap(); + if let Some(f) = handler.as_mut() { + f(data, is_binary); + } + } + + fn on_buffered_amount_change(&self, sent_data_size: u64) { + trace!("DataChannel: on_buffered_amount_change"); + let mut handler = self.on_buffered_amount_change_handler.lock().unwrap(); + if let Some(f) = handler.as_mut() { + f(sent_data_size); + } + } +} + +impl Default for InternalDataChannelObserver { + fn default() -> Self { + Self { + on_state_change_handler: Arc::new(Default::default()), + on_message_handler: Arc::new(Default::default()), + on_buffered_amount_change_handler: Arc::new(Default::default()), + } } } diff --git a/crates/livekit-webrtc/src/peer_connection.rs b/crates/livekit-webrtc/src/peer_connection.rs index 3789b3a..4329e2a 100644 --- a/crates/livekit-webrtc/src/peer_connection.rs +++ b/crates/livekit-webrtc/src/peer_connection.rs @@ -1,10 +1,9 @@ +use std::fmt::{Debug, Formatter}; use cxx::UniquePtr; use libwebrtc_sys::data_channel as sys_dc; use libwebrtc_sys::jsep as sys_jsep; use libwebrtc_sys::peer_connection as sys_pc; use log::trace; -use std::future::Future; -use std::pin::Pin; use std::sync::{Arc, Mutex}; use thiserror::Error; use tokio::sync::{mpsc, oneshot}; @@ -160,8 +159,11 @@ impl PeerConnection { tx.blocking_send(error).unwrap(); })); - let mut native_observer = sys_pc::ffi::create_native_add_ice_candidate_observer(Box::new(observer)); - self.cxx_handle.pin_mut().add_ice_candidate(candidate.release(), native_observer.pin_mut()); + let mut native_observer = + sys_pc::ffi::create_native_add_ice_candidate_observer(Box::new(observer)); + self.cxx_handle + .pin_mut() + .add_ice_candidate(candidate.release(), native_observer.pin_mut()); match rx.recv().await { Some(value) => Ok(()), @@ -446,7 +448,7 @@ impl sys_pc::PeerConnectionObserver for InternalObserver { trace!("on_data_channel"); let mut handler = self.on_data_channel_handler.lock().unwrap(); if let Some(f) = handler.as_mut() { - // TODO(theomonnom) + f(DataChannel::new(data_channel)); } } @@ -502,7 +504,7 @@ impl sys_pc::PeerConnectionObserver for InternalObserver { } fn on_ice_candidate(&self, candidate: UniquePtr) { - trace!("TESTING on_ice_candidate"); + trace!("on_ice_candidate"); let mut handler = self.on_ice_candidate_handler.lock().unwrap(); if let Some(f) = handler.as_mut() { f(IceCandidate::new(candidate)); @@ -602,11 +604,12 @@ impl sys_pc::PeerConnectionObserver for InternalObserver { #[cfg(test)] mod tests { - use crate::data_channel::DataChannelInit; + use crate::data_channel::{DataChannel, DataChannelInit}; use crate::jsep::IceCandidate; - use crate::peer_connection_factory::{PeerConnectionFactory, ICEServer, RTCConfiguration}; - use tokio::sync::mpsc; + use crate::peer_connection_factory::{ICEServer, PeerConnectionFactory, RTCConfiguration}; use crate::webrtc::RTCRuntime; + use log::trace; + use tokio::sync::mpsc; fn init_log() { let _ = env_logger::builder().is_test(true).try_init(); @@ -630,8 +633,9 @@ mod tests { let mut bob = factory.create_peer_connection(config.clone()).unwrap(); let mut alice = factory.create_peer_connection(config.clone()).unwrap(); - let (bob_ice_tx, mut bob_ice_rx) = mpsc::channel::(1); - let (alice_ice_tx, mut alice_ice_rx) = mpsc::channel::(1); + let (bob_ice_tx, mut bob_ice_rx) = mpsc::channel::(16); + let (alice_ice_tx, mut alice_ice_rx) = mpsc::channel::(16); + let (alice_dc_tx, mut alice_dc_rx) = mpsc::channel::(16); bob.on_ice_candidate(Box::new(move |candidate| { bob_ice_tx.blocking_send(candidate).unwrap(); @@ -641,13 +645,21 @@ mod tests { alice_ice_tx.blocking_send(candidate).unwrap(); })); - bob.create_data_channel("test_dc", DataChannelInit::default()) + alice.on_data_channel(Box::new(move |dc| { + alice_dc_tx.blocking_send(dc).unwrap(); + })); + + let mut bob_dc = bob + .create_data_channel("test_dc", DataChannelInit::default()) .unwrap(); let offer = bob.create_offer().await.unwrap(); + trace!("Bob offer: {:?}", offer); bob.set_local_description(offer.clone()).await.unwrap(); alice.set_remote_description(offer).await.unwrap(); + let answer = alice.create_answer().await.unwrap(); + trace!("Alice answer: {:?}", answer); alice.set_local_description(answer.clone()).await.unwrap(); bob.set_remote_description(answer).await.unwrap(); @@ -657,6 +669,15 @@ mod tests { bob.add_ice_candidate(alice_ice).await.unwrap(); alice.add_ice_candidate(bob_ice).await.unwrap(); + let (data_tx, mut data_rx) = mpsc::channel::(1); + let mut alice_dc = alice_dc_rx.recv().await.unwrap(); + alice_dc.on_message(Box::new(move |data, is_binary| { + data_tx.blocking_send(String::from_utf8_lossy(data).to_string()).unwrap(); + })); + + assert!(bob_dc.send(b"This is a test", true)); + assert_eq!(data_rx.recv().await.unwrap(), "This is a test"); + alice.close(); bob.close(); } diff --git a/crates/livekit-webrtc/src/webrtc.rs b/crates/livekit-webrtc/src/webrtc.rs index c267100..80484ae 100644 --- a/crates/livekit-webrtc/src/webrtc.rs +++ b/crates/livekit-webrtc/src/webrtc.rs @@ -2,13 +2,13 @@ use cxx::UniquePtr; use libwebrtc_sys::webrtc as sys_rtc; pub struct RTCRuntime { - cxx_handle: UniquePtr + cxx_handle: UniquePtr, } impl RTCRuntime { pub fn new() -> Self { Self { - cxx_handle: sys_rtc::ffi::create_rtc_runtime() + cxx_handle: sys_rtc::ffi::create_rtc_runtime(), } } -} \ No newline at end of file +}