fix: datachannel deadlocks & dispose crashes (#87)

- Fix datachannel deadlock on dispose
- Fix audio device being disposed on the wrong thread
- Allow multiple RtcRuntime to be created (Useful for unit tests where we use multiple PeerConnectionFactory)
This commit is contained in:
Théo Monnom
2023-06-10 18:49:00 +02:00
committed by GitHub
parent ebbd93b2d8
commit 33dfcb27b0
9 changed files with 90 additions and 39 deletions
+9 -13
View File
@@ -247,13 +247,9 @@ mod tests {
use log::trace; use log::trace;
use tokio::sync::mpsc; use tokio::sync::mpsc;
fn init_log() {
let _ = env_logger::builder().is_test(true).try_init();
}
#[tokio::test] #[tokio::test]
async fn create_pc() { async fn create_pc() {
init_log(); let _ = env_logger::builder().is_test(true).try_init();
let factory = PeerConnectionFactory::default(); let factory = PeerConnectionFactory::default();
let config = RtcConfiguration { let config = RtcConfiguration {
@@ -269,20 +265,20 @@ mod tests {
let bob = factory.create_peer_connection(config.clone()).unwrap(); let bob = factory.create_peer_connection(config.clone()).unwrap();
let alice = factory.create_peer_connection(config.clone()).unwrap(); let alice = factory.create_peer_connection(config.clone()).unwrap();
let (bob_ice_tx, mut bob_ice_rx) = mpsc::channel::<IceCandidate>(16); let (bob_ice_tx, mut bob_ice_rx) = mpsc::unbounded_channel::<IceCandidate>();
let (alice_ice_tx, mut alice_ice_rx) = mpsc::channel::<IceCandidate>(16); let (alice_ice_tx, mut alice_ice_rx) = mpsc::unbounded_channel::<IceCandidate>();
let (alice_dc_tx, mut alice_dc_rx) = mpsc::channel::<DataChannel>(16); let (alice_dc_tx, mut alice_dc_rx) = mpsc::unbounded_channel::<DataChannel>();
bob.on_ice_candidate(Some(Box::new(move |candidate| { bob.on_ice_candidate(Some(Box::new(move |candidate| {
bob_ice_tx.blocking_send(candidate).unwrap(); bob_ice_tx.send(candidate).unwrap();
}))); })));
alice.on_ice_candidate(Some(Box::new(move |candidate| { alice.on_ice_candidate(Some(Box::new(move |candidate| {
alice_ice_tx.blocking_send(candidate).unwrap(); alice_ice_tx.send(candidate).unwrap();
}))); })));
alice.on_data_channel(Some(Box::new(move |dc| { alice.on_data_channel(Some(Box::new(move |dc| {
alice_dc_tx.blocking_send(dc).unwrap(); alice_dc_tx.send(dc).unwrap();
}))); })));
let bob_dc = bob let bob_dc = bob
@@ -305,11 +301,11 @@ mod tests {
bob.add_ice_candidate(alice_ice).await.unwrap(); bob.add_ice_candidate(alice_ice).await.unwrap();
alice.add_ice_candidate(bob_ice).await.unwrap(); alice.add_ice_candidate(bob_ice).await.unwrap();
let (data_tx, mut data_rx) = mpsc::channel::<String>(1); let (data_tx, mut data_rx) = mpsc::unbounded_channel::<String>();
let alice_dc = alice_dc_rx.recv().await.unwrap(); let alice_dc = alice_dc_rx.recv().await.unwrap();
alice_dc.on_message(Some(Box::new(move |buffer| { alice_dc.on_message(Some(Box::new(move |buffer| {
data_tx data_tx
.blocking_send(String::from_utf8_lossy(buffer.data).to_string()) .send(String::from_utf8_lossy(buffer.data).to_string())
.unwrap(); .unwrap();
}))); })));
+1 -2
View File
@@ -40,6 +40,7 @@ class DataChannel {
explicit DataChannel( explicit DataChannel(
std::shared_ptr<RtcRuntime> rtc_runtime, std::shared_ptr<RtcRuntime> rtc_runtime,
rtc::scoped_refptr<webrtc::DataChannelInterface> data_channel); rtc::scoped_refptr<webrtc::DataChannelInterface> data_channel);
~DataChannel();
void register_observer(rust::Box<DataChannelObserverWrapper> observer) const; void register_observer(rust::Box<DataChannelObserverWrapper> observer) const;
void unregister_observer() const; void unregister_observer() const;
@@ -64,8 +65,6 @@ class NativeDataChannelObserver : public webrtc::DataChannelObserver {
NativeDataChannelObserver(rust::Box<DataChannelObserverWrapper> observer, NativeDataChannelObserver(rust::Box<DataChannelObserverWrapper> observer,
const DataChannel* dc); const DataChannel* dc);
~NativeDataChannelObserver();
void OnStateChange() override; void OnStateChange() override;
void OnMessage(const webrtc::DataBuffer& buffer) override; void OnMessage(const webrtc::DataBuffer& buffer) override;
void OnBufferedAmountChange(uint64_t sent_data_size) override; void OnBufferedAmountChange(uint64_t sent_data_size) override;
@@ -48,6 +48,8 @@ class PeerConnection {
std::unique_ptr<NativePeerConnectionObserver> observer, std::unique_ptr<NativePeerConnectionObserver> observer,
rtc::scoped_refptr<webrtc::PeerConnectionInterface> peer_connection); rtc::scoped_refptr<webrtc::PeerConnectionInterface> peer_connection);
~PeerConnection();
void create_offer( void create_offer(
RtcOfferAnswerOptions options, RtcOfferAnswerOptions options,
rust::Box<AsyncContext> ctx, rust::Box<AsyncContext> ctx,
@@ -17,6 +17,8 @@
#pragma once #pragma once
#include "api/peer_connection_interface.h" #include "api/peer_connection_interface.h"
#include "api/scoped_refptr.h"
#include "livekit/audio_device.h"
#include "media_stream.h" #include "media_stream.h"
#include "peer_connection.h" #include "peer_connection.h"
#include "rtp_parameters.h" #include "rtp_parameters.h"
@@ -57,6 +59,7 @@ class PeerConnectionFactory {
private: private:
std::shared_ptr<RtcRuntime> rtc_runtime_; std::shared_ptr<RtcRuntime> rtc_runtime_;
rtc::scoped_refptr<AudioDevice> audio_device_;
rtc::scoped_refptr<webrtc::PeerConnectionFactoryInterface> peer_factory_; rtc::scoped_refptr<webrtc::PeerConnectionFactoryInterface> peer_factory_;
}; };
+5 -3
View File
@@ -80,6 +80,8 @@ class RtcRuntime : public std::enable_shared_from_this<RtcRuntime> {
// have one livekit::VideoTrack associated with it). // have one livekit::VideoTrack associated with it).
// The only reason we to do that is to allow to add states inside our // The only reason we to do that is to allow to add states inside our
// wrappers (e.g: the sinks_ member inside AudioTrack) // wrappers (e.g: the sinks_ member inside AudioTrack)
// DataChannel and the PeerConnectionFactory don't need to do this (There's no
// way to retrieve them after creation)
webrtc::Mutex mutex_; webrtc::Mutex mutex_;
std::vector<std::weak_ptr<MediaStreamTrack>> media_stream_tracks_; std::vector<std::weak_ptr<MediaStreamTrack>> media_stream_tracks_;
// We don't have additonal state in RtpReceiver and RtpSender atm.. // We don't have additonal state in RtpReceiver and RtpSender atm..
@@ -87,9 +89,9 @@ class RtcRuntime : public std::enable_shared_from_this<RtcRuntime> {
// std::vector<std::weak_ptr<RtpSender>> rtp_senders_; // std::vector<std::weak_ptr<RtpSender>> rtp_senders_;
#ifdef WEBRTC_WIN #ifdef WEBRTC_WIN
rtc::WinsockInitializer winsock_; // rtc::WinsockInitializer winsock_;
rtc::PhysicalSocketServer ss_; // rtc::PhysicalSocketServer ss_;
rtc::AutoSocketServerThread main_thread_{&ss_}; // rtc::AutoSocketServerThread main_thread_{&ss_};
#endif #endif
}; };
+4 -4
View File
@@ -47,6 +47,10 @@ DataChannel::DataChannel(
rtc::scoped_refptr<webrtc::DataChannelInterface> data_channel) rtc::scoped_refptr<webrtc::DataChannelInterface> data_channel)
: rtc_runtime_(rtc_runtime), data_channel_(std::move(data_channel)) {} : rtc_runtime_(rtc_runtime), data_channel_(std::move(data_channel)) {}
DataChannel::~DataChannel() {
unregister_observer();
}
void DataChannel::register_observer( void DataChannel::register_observer(
rust::Box<DataChannelObserverWrapper> observer) const { rust::Box<DataChannelObserverWrapper> observer) const {
webrtc::MutexLock lock(&mutex_); webrtc::MutexLock lock(&mutex_);
@@ -86,10 +90,6 @@ NativeDataChannelObserver::NativeDataChannelObserver(
const DataChannel* dc) const DataChannel* dc)
: observer_(std::move(observer)), dc_(dc) {} : observer_(std::move(observer)), dc_(dc) {}
NativeDataChannelObserver::~NativeDataChannelObserver() {
dc_->unregister_observer();
}
void NativeDataChannelObserver::OnStateChange() { void NativeDataChannelObserver::OnStateChange() {
observer_->on_state_change(dc_->state()); observer_->on_state_change(dc_->state());
} }
+8 -1
View File
@@ -25,6 +25,7 @@
#include "livekit/media_stream.h" #include "livekit/media_stream.h"
#include "livekit/rtc_error.h" #include "livekit/rtc_error.h"
#include "livekit/rtp_transceiver.h" #include "livekit/rtp_transceiver.h"
#include "rtc_base/logging.h"
#include "webrtc-sys/src/peer_connection.rs.h" #include "webrtc-sys/src/peer_connection.rs.h"
#include "webrtc-sys/src/rtc_error.rs.h" #include "webrtc-sys/src/rtc_error.rs.h"
@@ -50,7 +51,13 @@ PeerConnection::PeerConnection(
rtc::scoped_refptr<webrtc::PeerConnectionInterface> peer_connection) rtc::scoped_refptr<webrtc::PeerConnectionInterface> peer_connection)
: rtc_runtime_(rtc_runtime), : rtc_runtime_(rtc_runtime),
observer_(std::move(observer)), observer_(std::move(observer)),
peer_connection_(std::move(peer_connection)) {} peer_connection_(std::move(peer_connection)) {
RTC_LOG(LS_VERBOSE) << "PeerConnection::PeerConnection()";
}
PeerConnection::~PeerConnection() {
RTC_LOG(LS_VERBOSE) << "PeerConnection::~PeerConnection()";
}
void PeerConnection::create_offer( void PeerConnection::create_offer(
RtcOfferAnswerOptions options, RtcOfferAnswerOptions options,
+14 -8
View File
@@ -70,7 +70,7 @@ webrtc::PeerConnectionInterface::RTCConfiguration to_native_rtc_configuration(
PeerConnectionFactory::PeerConnectionFactory( PeerConnectionFactory::PeerConnectionFactory(
std::shared_ptr<RtcRuntime> rtc_runtime) std::shared_ptr<RtcRuntime> rtc_runtime)
: rtc_runtime_(rtc_runtime) { : rtc_runtime_(rtc_runtime) {
RTC_LOG(LS_INFO) << "PeerConnectionFactory::PeerConnectionFactory()"; RTC_LOG(LS_VERBOSE) << "PeerConnectionFactory::PeerConnectionFactory()";
webrtc::PeerConnectionFactoryDependencies dependencies; webrtc::PeerConnectionFactoryDependencies dependencies;
dependencies.network_thread = rtc_runtime_->network_thread(); dependencies.network_thread = rtc_runtime_->network_thread();
@@ -86,12 +86,14 @@ PeerConnectionFactory::PeerConnectionFactory(
cricket::MediaEngineDependencies media_deps; cricket::MediaEngineDependencies media_deps;
media_deps.task_queue_factory = dependencies.task_queue_factory.get(); media_deps.task_queue_factory = dependencies.task_queue_factory.get();
media_deps.adm = rtc_runtime_->worker_thread() audio_device_ = rtc_runtime_->worker_thread()
->Invoke<rtc::scoped_refptr<livekit::AudioDevice>>( ->Invoke<rtc::scoped_refptr<livekit::AudioDevice>>(
RTC_FROM_HERE, [&] { RTC_FROM_HERE, [&] {
return rtc::make_ref_counted<livekit::AudioDevice>( return rtc::make_ref_counted<livekit::AudioDevice>(
media_deps.task_queue_factory); media_deps.task_queue_factory);
}); });
media_deps.adm = audio_device_;
media_deps.video_encoder_factory = media_deps.video_encoder_factory =
std::move(std::make_unique<livekit::VideoEncoderFactory>()); std::move(std::make_unique<livekit::VideoEncoderFactory>());
@@ -114,7 +116,11 @@ PeerConnectionFactory::PeerConnectionFactory(
} }
PeerConnectionFactory::~PeerConnectionFactory() { PeerConnectionFactory::~PeerConnectionFactory() {
RTC_LOG(LS_INFO) << "PeerConnectionFactory::~PeerConnectionFactory()"; RTC_LOG(LS_VERBOSE) << "PeerConnectionFactory::~PeerConnectionFactory()";
peer_factory_ = nullptr;
rtc_runtime_->worker_thread()->Invoke<void>(
RTC_FROM_HERE, [this] { audio_device_ = nullptr; });
} }
std::shared_ptr<PeerConnection> PeerConnectionFactory::create_peer_connection( std::shared_ptr<PeerConnection> PeerConnectionFactory::create_peer_connection(
+44 -8
View File
@@ -16,6 +16,9 @@
#include "livekit/webrtc.h" #include "livekit/webrtc.h"
#include <algorithm>
#include <atomic>
#include <iostream>
#include <memory> #include <memory>
#include "livekit/audio_track.h" #include "livekit/audio_track.h"
@@ -27,10 +30,34 @@
#include "rtc_base/logging.h" #include "rtc_base/logging.h"
#include "rtc_base/synchronization/mutex.h" #include "rtc_base/synchronization/mutex.h"
#ifdef WEBRTC_WIN
#include "rtc_base/win32.h"
#endif
namespace livekit { namespace livekit {
static webrtc::Mutex g_mutex{};
// Can't be atomic, we're using a Mutex because we need to wait for the
// execution of the first init
static uint32_t g_release_counter(0);
RtcRuntime::RtcRuntime() { RtcRuntime::RtcRuntime() {
RTC_LOG(LS_INFO) << "RtcRuntime()"; rtc::LogMessage::LogToDebug(rtc::LS_INFO);
RTC_CHECK(rtc::InitializeSSL()) << "Failed to InitializeSSL()"; RTC_LOG(LS_VERBOSE) << "RtcRuntime()";
{
// Not the best way to do it...
webrtc::MutexLock lock(&g_mutex);
if (g_release_counter == 0) {
RTC_CHECK(rtc::InitializeSSL()) << "Failed to InitializeSSL()";
#ifdef WEBRTC_WIN
WSADATA data;
WSAStartup(MAKEWORD(1, 0), &data);
#endif
}
g_release_counter++;
}
network_thread_ = rtc::Thread::CreateWithSocketServer(); network_thread_ = rtc::Thread::CreateWithSocketServer();
network_thread_->SetName("network_thread", &network_thread_); network_thread_->SetName("network_thread", &network_thread_);
@@ -44,14 +71,23 @@ RtcRuntime::RtcRuntime() {
} }
RtcRuntime::~RtcRuntime() { RtcRuntime::~RtcRuntime() {
RTC_LOG(LS_INFO) << "~RtcRuntime()"; RTC_LOG(LS_VERBOSE) << "~RtcRuntime()";
rtc::ThreadManager::Instance()->SetCurrentThread(nullptr); worker_thread_->Stop();
RTC_CHECK(rtc::CleanupSSL()) << "Failed to CleanupSSL()"; signaling_thread_->Stop();
network_thread_->Stop();
worker_thread_->Quit(); {
signaling_thread_->Quit(); webrtc::MutexLock lock(&g_mutex);
network_thread_->Quit(); g_release_counter--;
if (g_release_counter == 0) {
RTC_CHECK(rtc::CleanupSSL()) << "Failed to CleanupSSL()";
#ifdef WEBRTC_WIN
WSACleanup();
#endif
}
}
} }
rtc::Thread* RtcRuntime::network_thread() const { rtc::Thread* RtcRuntime::network_thread() const {