Move signaling safety flag into SctpDataChannel and clarify its purpose Encapsulate the signaling thread safety flag within SctpDataChannel and rename to `controller_safety` to clarify its purpose. Remove details about the flag from the CreateProxy() interface. The observer previously received this flag via the CreateProxy method, but the observer is associated with the channel (not the proxy). Basically, the proxy creation layer should not need to know about the flag mechanism and the channel object was out of the loop, so this was a bit awkward. Bug: webrtc:510487699 Change-Id: I57638ade52305b6479a9a9165747336c679631cd Reviewed-on: https://webrtc-review.googlesource.com/c/src/+/304841 Reviewed-by: Danil Chapovalov <danilchap@webrtc.org> Commit-Queue: Tomas Gunnarsson <tommi@webrtc.org> Cr-Commit-Position: refs/heads/main@{#47649}
diff --git a/pc/data_channel_controller.cc b/pc/data_channel_controller.cc index c8a4e1a..3ce3c26 100644 --- a/pc/data_channel_controller.cc +++ b/pc/data_channel_controller.cc
@@ -343,7 +343,7 @@ scoped_refptr<SctpDataChannel> channel, bool ready_to_send) { channel_usage_ = DataChannelUsage::kInUse; - auto proxy = SctpDataChannel::CreateProxy(channel, signaling_safety_.flag()); + auto proxy = SctpDataChannel::CreateProxy(channel); pc_->RunWithObserver([&](auto observer) { observer->OnDataChannel(proxy); }); pc_->NoteDataAddedEvent(); @@ -408,7 +408,7 @@ scoped_refptr<SctpDataChannel> channel = SctpDataChannel::Create( weak_factory_.GetWeakPtr(), label, data_channel_transport_ != nullptr, - config, signaling_thread(), network_thread()); + config, signaling_safety_.flag(), signaling_thread(), network_thread()); RTC_DCHECK(channel); // If we have an id already, notify the transport. @@ -465,8 +465,7 @@ return ret.MoveError(); channel_usage_ = DataChannelUsage::kInUse; - return SctpDataChannel::CreateProxy(ret.MoveValue(), - signaling_safety_.flag()); + return SctpDataChannel::CreateProxy(ret.MoveValue()); } void DataChannelController::AllocateSctpSids(SSLRole role) {
diff --git a/pc/rtc_stats_collector_unittest.cc b/pc/rtc_stats_collector_unittest.cc index 57f7e84..9883293 100644 --- a/pc/rtc_stats_collector_unittest.cc +++ b/pc/rtc_stats_collector_unittest.cc
@@ -2201,13 +2201,15 @@ EXPECT_EQ(expected, report->Get("P")->cast_to<RTCPeerConnectionStats>()); } + ScopedTaskSafety signaling_safety; + FakeDataChannelController controller(pc_->network_thread()); scoped_refptr<SctpDataChannel> dummy_channel_a = SctpDataChannel::Create( controller.weak_ptr(), "DummyChannelA", false, InternalDataChannelInit(), - Thread::Current(), Thread::Current()); + signaling_safety.flag(), Thread::Current(), Thread::Current()); scoped_refptr<SctpDataChannel> dummy_channel_b = SctpDataChannel::Create( controller.weak_ptr(), "DummyChannelB", false, InternalDataChannelInit(), - Thread::Current(), Thread::Current()); + signaling_safety.flag(), Thread::Current(), Thread::Current()); stats_->stats_collector().OnSctpDataChannelStateChanged( dummy_channel_a->internal_id(), DataChannelInterface::DataState::kOpen);
diff --git a/pc/sctp_data_channel.cc b/pc/sctp_data_channel.cc index 6d7305a..6aab956 100644 --- a/pc/sctp_data_channel.cc +++ b/pc/sctp_data_channel.cc
@@ -179,10 +179,7 @@ // and the ObserverAdapter no longer be necessary. class SctpDataChannel::ObserverAdapter : public DataChannelObserver { public: - explicit ObserverAdapter( - SctpDataChannel* channel, - scoped_refptr<PendingTaskSafetyFlag> signaling_safety) - : channel_(channel), signaling_safety_(std::move(signaling_safety)) {} + explicit ObserverAdapter(SctpDataChannel* channel) : channel_(channel) {} bool IsInsideCallback() const { RTC_DCHECK_RUN_ON(signaling_thread()); @@ -237,7 +234,8 @@ RTC_DCHECK(was_dropped_); was_dropped_ = false; adapter_->cached_getters_ = this; - return adapter_->delegate_ && adapter_->signaling_safety_->alive(); + return adapter_->delegate_ && + adapter_->channel_->controller_safety_->alive(); } RTCError error() { return cached_error_; } @@ -294,7 +292,6 @@ // `channel_` in the `RTC_DCHECK_RUN_ON` checks on the signaling thread. Thread* const signaling_thread_{channel_->signaling_thread_}; ScopedTaskSafety safety_; - scoped_refptr<PendingTaskSafetyFlag> signaling_safety_; CachedGetters* cached_getters_ RTC_GUARDED_BY(signaling_thread()) = nullptr; }; @@ -304,23 +301,22 @@ absl::string_view label, bool connected_to_transport, const InternalDataChannelInit& config, + scoped_refptr<PendingTaskSafetyFlag> controller_safety, Thread* signaling_thread, Thread* network_thread) { RTC_DCHECK(config.IsValid()); - return make_ref_counted<SctpDataChannel>(config, std::move(controller), label, - connected_to_transport, - signaling_thread, network_thread); + return make_ref_counted<SctpDataChannel>( + config, std::move(controller), label, connected_to_transport, + std::move(controller_safety), signaling_thread, network_thread); } // static scoped_refptr<DataChannelInterface> SctpDataChannel::CreateProxy( - scoped_refptr<SctpDataChannel> channel, - scoped_refptr<PendingTaskSafetyFlag> signaling_safety) { + scoped_refptr<SctpDataChannel> channel) { // Copy thread params to local variables before `std::move()`. auto* signaling_thread = channel->signaling_thread_; auto* network_thread = channel->network_thread_; - channel->observer_adapter_ = std::make_unique<ObserverAdapter>( - channel.get(), std::move(signaling_safety)); + channel->observer_adapter_ = std::make_unique<ObserverAdapter>(channel.get()); return DataChannelProxy::Create(signaling_thread, network_thread, std::move(channel)); } @@ -330,6 +326,7 @@ WeakPtr<SctpDataChannelControllerInterface> controller, absl::string_view label, bool connected_to_transport, + scoped_refptr<PendingTaskSafetyFlag> controller_safety, Thread* signaling_thread, Thread* network_thread) : signaling_thread_(signaling_thread), @@ -345,7 +342,8 @@ ordered_(config.ordered), observer_(nullptr), controller_(std::move(controller)), - connected_to_transport_(connected_to_transport) { + connected_to_transport_(connected_to_transport), + controller_safety_(std::move(controller_safety)) { RTC_DCHECK_RUN_ON(network_thread_); // Since we constructed on the network thread we can't (yet) check the // `controller_` pointer since doing so will trigger a thread check.
diff --git a/pc/sctp_data_channel.h b/pc/sctp_data_channel.h index 693661b..7028022 100644 --- a/pc/sctp_data_channel.h +++ b/pc/sctp_data_channel.h
@@ -140,24 +140,24 @@ // OnClosingProcedureComplete callback and transition to kClosed. class SctpDataChannel : public DataChannelInterface { public: + // The `controller_safety` flag is used for the ObserverAdapter callback proxy + // which delivers callbacks on the `signaling_thread` but must not deliver + // such callbacks after the peerconnection has been closed. The data + // controller will update the flag when closed, which will cancel any pending + // event notifications. static scoped_refptr<SctpDataChannel> Create( WeakPtr<SctpDataChannelControllerInterface> controller, absl::string_view label, bool connected_to_transport, const InternalDataChannelInit& config, + scoped_refptr<PendingTaskSafetyFlag> controller_safety, Thread* signaling_thread, Thread* network_thread); // Instantiates an API proxy for a SctpDataChannel instance that will be // handed out to external callers. - // The `signaling_safety` flag is used for the ObserverAdapter callback proxy - // which delivers callbacks on the signaling thread but must not deliver such - // callbacks after the peerconnection has been closed. The data controller - // will update the flag when closed, which will cancel any pending event - // notifications. static scoped_refptr<DataChannelInterface> CreateProxy( - scoped_refptr<SctpDataChannel> channel, - scoped_refptr<PendingTaskSafetyFlag> signaling_safety); + scoped_refptr<SctpDataChannel> channel); void RegisterObserver(DataChannelObserver* observer) override; void UnregisterObserver() override; @@ -244,6 +244,7 @@ WeakPtr<SctpDataChannelControllerInterface> controller, absl::string_view label, bool connected_to_transport, + scoped_refptr<PendingTaskSafetyFlag> controller_safety, Thread* signaling_thread, Thread* network_thread); ~SctpDataChannel() override; @@ -308,6 +309,7 @@ bool started_closing_procedure_ RTC_GUARDED_BY(network_thread_) = false; bool connected_to_transport_ RTC_GUARDED_BY(network_thread_) = false; PacketQueue queued_received_data_ RTC_GUARDED_BY(network_thread_); + scoped_refptr<PendingTaskSafetyFlag> controller_safety_; }; } // namespace webrtc
diff --git a/pc/sctp_data_channel_unittest.cc b/pc/sctp_data_channel_unittest.cc index c4e30ee..1223f9e 100644 --- a/pc/sctp_data_channel_unittest.cc +++ b/pc/sctp_data_channel_unittest.cc
@@ -24,7 +24,6 @@ #include "api/scoped_refptr.h" #include "api/sctp_transport_interface.h" #include "api/sequence_checker.h" -#include "api/task_queue/pending_task_safety_flag.h" #include "api/test/rtc_error_matchers.h" #include "api/transport/data_channel_transport_interface.h" #include "pc/sctp_utils.h" @@ -88,11 +87,10 @@ controller_(new FakeDataChannelController(&network_thread_)) { network_thread_.Start(); inner_channel_ = controller_->CreateDataChannel("test", init_); - channel_ = SctpDataChannel::CreateProxy(inner_channel_, signaling_safety_); + channel_ = SctpDataChannel::CreateProxy(inner_channel_); } ~SctpDataChannelTest() override { run_loop_.Flush(); - signaling_safety_->SetNotAlive(); inner_channel_ = nullptr; channel_ = nullptr; controller_.reset(); @@ -150,8 +148,6 @@ test::RunLoop run_loop_; Thread network_thread_; InternalDataChannelInit init_; - scoped_refptr<PendingTaskSafetyFlag> signaling_safety_ = - PendingTaskSafetyFlag::Create(); std::unique_ptr<FakeDataChannelController> controller_; std::unique_ptr<FakeDataChannelObserver> observer_; scoped_refptr<SctpDataChannel> inner_channel_; @@ -249,7 +245,7 @@ InternalDataChannelInit init; init.id = 1; auto dc = SctpDataChannel::CreateProxy( - controller_->CreateDataChannel("test1", init), signaling_safety_); + controller_->CreateDataChannel("test1", init)); EXPECT_EQ(DataChannelInterface::kOpen, dc->state()); } @@ -262,7 +258,7 @@ init.ordered = false; scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", init); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)), @@ -293,7 +289,7 @@ init.ordered = false; scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", init); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)), @@ -324,7 +320,7 @@ init.ordered = false; scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", init); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)), @@ -349,7 +345,7 @@ init.ordered = false; scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", init); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)), @@ -410,7 +406,7 @@ SetChannelReady(); scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", config); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)), @@ -477,7 +473,7 @@ SetChannelReady(); scoped_refptr<SctpDataChannel> dc = controller_->CreateDataChannel("test1", config); - auto proxy = SctpDataChannel::CreateProxy(dc, signaling_safety_); + auto proxy = SctpDataChannel::CreateProxy(dc); EXPECT_THAT(WaitUntil([&] { return proxy->state(); }, Eq(DataChannelInterface::kOpen)),
diff --git a/pc/test/fake_data_channel_controller.h b/pc/test/fake_data_channel_controller.h index f6368a5..c003f4a 100644 --- a/pc/test/fake_data_channel_controller.h +++ b/pc/test/fake_data_channel_controller.h
@@ -23,6 +23,7 @@ #include "api/rtc_error.h" #include "api/scoped_refptr.h" #include "api/sequence_checker.h" +#include "api/task_queue/pending_task_safety_flag.h" #include "api/transport/data_channel_transport_interface.h" #include "pc/sctp_data_channel.h" #include "pc/sctp_utils.h" @@ -70,7 +71,8 @@ scoped_refptr<SctpDataChannel> channel = SctpDataChannel::Create( std::move(my_weak_ptr), std::string(label), transport_available_, - init, signaling_thread_, network_thread_); + init, signaling_safety_.flag(), signaling_thread_, + network_thread_); if (transport_available_ && channel->sid_n().has_value()) { AddSctpDataStream(*channel->sid_n(), channel->priority()); } @@ -236,6 +238,7 @@ private: Thread* const signaling_thread_; Thread* const network_thread_; + ScopedTaskSafety signaling_safety_; StreamId last_sid_ RTC_GUARDED_BY(network_thread_); SendDataParams last_send_data_params_ RTC_GUARDED_BY(network_thread_); bool send_blocked_ RTC_GUARDED_BY(network_thread_);
diff --git a/pc/test/fake_peer_connection_for_stats.h b/pc/test/fake_peer_connection_for_stats.h index 0de9786..18fce85 100644 --- a/pc/test/fake_peer_connection_for_stats.h +++ b/pc/test/fake_peer_connection_for_stats.h
@@ -505,10 +505,9 @@ void AddSctpDataChannel(const std::string& label, const InternalDataChannelInit& init) { - // TODO(bugs.webrtc.org/11547): Supply a separate network thread. AddSctpDataChannel(SctpDataChannel::Create( data_channel_controller_.weak_ptr(), label, false, init, - Thread::Current(), Thread::Current())); + signaling_safety_.flag(), signaling_thread_, network_thread_)); } void AddSctpDataChannel(scoped_refptr<SctpDataChannel> data_channel) { @@ -770,6 +769,7 @@ Thread* const network_thread_; Thread* const worker_thread_; Thread* const signaling_thread_; + ScopedTaskSafety signaling_safety_; PeerConnectionFactoryDependencies dependencies_; scoped_refptr<ConnectionContext> context_;
diff --git a/pc/test/mock_data_channel.h b/pc/test/mock_data_channel.h index f2822cb..fdc554c 100644 --- a/pc/test/mock_data_channel.h +++ b/pc/test/mock_data_channel.h
@@ -15,6 +15,7 @@ #include <string> #include <utility> +#include "api/task_queue/pending_task_safety_flag.h" #include "pc/sctp_data_channel.h" #include "rtc_base/thread.h" #include "rtc_base/weak_ptr.h" @@ -53,6 +54,7 @@ std::move(controller), label, false, + PendingTaskSafetyFlag::Create(), signaling_thread, network_thread) { EXPECT_CALL(*this, id()).WillRepeatedly(::testing::Return(id));