diff --git a/cpp/core/internal/mediums/BUILD b/cpp/core/internal/mediums/BUILD index 8106972d..6b9fb030 100644 --- a/cpp/core/internal/mediums/BUILD +++ b/cpp/core/internal/mediums/BUILD @@ -52,6 +52,7 @@ cc_library( "//location/nearby/mediums/proto:web_rtc_signaling_frames_cc_proto", "//absl/container:flat_hash_map", "//absl/container:flat_hash_set", + "//absl/functional:bind_front", "//absl/numeric:int128", "//absl/strings", "//absl/time", diff --git a/cpp/core/internal/mediums/webrtc.cc b/cpp/core/internal/mediums/webrtc.cc index 3323b5e6..4caeca0f 100644 --- a/cpp/core/internal/mediums/webrtc.cc +++ b/cpp/core/internal/mediums/webrtc.cc @@ -27,6 +27,7 @@ #include "platform/public/logging.h" #include "platform/public/mutex_lock.h" #include "location/nearby/mediums/proto/web_rtc_signaling_frames.pb.h" +#include "absl/functional/bind_front.h" #include "absl/strings/str_cat.h" #include "absl/time/time.h" #include "webrtc/api/jsep.h" @@ -122,14 +123,9 @@ bool WebRtc::StartAcceptingConnections(const std::string& service_id, // This registers ourselves w/ Tachyon, creating a room from the PeerId. // This allows a remote device to message us over Tachyon. - auto signaling_message_callback = [this, service_id](ByteArray message) { - OffloadFromThread([this, service_id{std::move(service_id)}, - message{std::move(message)}]() { - ProcessTachyonInboxMessage(service_id, message); - }); - }; if (!info.signaling_messenger->StartReceivingMessages( - signaling_message_callback)) { + absl::bind_front(&WebRtc::OnSignalingMessage, this, service_id), + absl::bind_front(&WebRtc::OnSignalingComplete, this, service_id))) { info.signaling_messenger.reset(); return false; } @@ -270,14 +266,16 @@ WebRtcSocketWrapper WebRtc::AttemptToConnect( // This registers ourselves w/ Tachyon, creating a room from the PeerId. // This allows a remote device to message us over Tachyon. - auto signaling_message_callback = [this, service_id](ByteArray message) { - OffloadFromThread([this, service_id{std::move(service_id)}, - message{std::move(message)}]() { - ProcessTachyonInboxMessage(service_id, message); - }); + auto signaling_complete_callback = [this, &socket_future](bool success) { + if (!success) { + OffloadFromThread([&socket_future]() { + socket_future.SetException({Exception::kFailed}); + }); + } }; if (!info.signaling_messenger->StartReceivingMessages( - signaling_message_callback)) { + absl::bind_front(&WebRtc::OnSignalingMessage, this, service_id), + signaling_complete_callback)) { NEARBY_LOG(INFO, "Cannot connect to WebRTC peer %s because we failed to start " "receiving messages over Tachyon.", @@ -392,6 +390,37 @@ void WebRtc::ProcessLocalIceCandidate( service_id.c_str()); } +void WebRtc::OnSignalingMessage(const std::string& service_id, + const ByteArray& message) { + OffloadFromThread([this, service_id, message]() { + ProcessTachyonInboxMessage(service_id, message); + }); +} + +void WebRtc::OnSignalingComplete(const std::string& service_id, bool success) { + NEARBY_LOG(INFO, "Signaling completed with status: %d.", success); + if (success) { + return; + } + + OffloadFromThread([this, service_id]() { + MutexLock lock(&mutex_); + const auto& info_entry = accepting_connections_info_.find(service_id); + if (info_entry == accepting_connections_info_.end()) { + return; + } + + if (info_entry->second.restart_accept_connections_count < + kRestartAcceptConnectionsLimit) { + ++info_entry->second.restart_accept_connections_count; + } else { + return; + } + + RestartTachyonReceiveMessages(service_id); + }); +} + void WebRtc::ProcessTachyonInboxMessage(const std::string& service_id, const ByteArray& message) { MutexLock lock(&mutex_); @@ -593,6 +622,10 @@ void WebRtc::ReceiveIceCandidates( void WebRtc::ProcessRestartTachyonReceiveMessages( const std::string& service_id) { MutexLock lock(&mutex_); + RestartTachyonReceiveMessages(service_id); +} + +void WebRtc::RestartTachyonReceiveMessages(const std::string& service_id) { if (!IsAcceptingConnectionsLocked(service_id)) { NEARBY_LOG(INFO, "Skipping restart listening for tachyon inbox messages since we " @@ -608,14 +641,9 @@ void WebRtc::ProcessRestartTachyonReceiveMessages( info.signaling_messenger->StopReceivingMessages(); // Attempt to re-register. - auto signaling_message_callback = [this, service_id](ByteArray message) { - OffloadFromThread([this, service_id{std::move(service_id)}, - message{std::move(message)}]() { - ProcessTachyonInboxMessage(service_id, message); - }); - }; if (!info.signaling_messenger->StartReceivingMessages( - signaling_message_callback)) { + absl::bind_front(&WebRtc::OnSignalingMessage, this, service_id), + absl::bind_front(&WebRtc::OnSignalingComplete, this, service_id))) { NEARBY_LOG(WARNING, "Failed to restart listening for tachyon inbox messages for " "service %s since we failed to reach Tachyon.", diff --git a/cpp/core/internal/mediums/webrtc.h b/cpp/core/internal/mediums/webrtc.h index 06541980..ac3ae2c2 100644 --- a/cpp/core/internal/mediums/webrtc.h +++ b/cpp/core/internal/mediums/webrtc.h @@ -15,6 +15,7 @@ #ifndef CORE_INTERNAL_MEDIUMS_WEBRTC_H_ #define CORE_INTERNAL_MEDIUMS_WEBRTC_H_ +#include #include #include @@ -98,6 +99,9 @@ class WebRtc { ABSL_LOCKS_EXCLUDED(mutex_); private: + static constexpr int kConnectAttemptsLimit = 3; + static constexpr int kRestartAcceptConnectionsLimit = 3; + enum class Role { kNone = 0, kOfferer = 1, @@ -121,6 +125,11 @@ class WebRtc { // advertising. Non-null when listening for WebRTC connections as an // offerer. CancelableAlarm restart_tachyon_receive_messages_alarm; + + // Tracks the number of times we've restarted receiving messages after a + // failure. We limit the number to prevent endless restarts if we are + // repeatedly unable to communicate with Tachyon. + int restart_accept_connections_count = 0; }; struct ConnectionRequestInfo { @@ -136,8 +145,6 @@ class WebRtc { Future socket_future; }; - static constexpr int kConnectAttemptsLimit = 3; - // Attempt to initiates a WebRtc connection with peer device identified by // |peer_id|. // Runs on @MainThread. @@ -152,6 +159,13 @@ class WebRtc { bool IsAcceptingConnectionsLocked(const std::string& service_id) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + // Receives a message from the signaling messenger. + void OnSignalingMessage(const std::string& service_id, + const ByteArray& message); + + // Decides whether to restart receiving messages. + void OnSignalingComplete(const std::string& service_id, bool success); + // Runs on |single_thread_executor_|. void ProcessTachyonInboxMessage(const std::string& service_id, const ByteArray& message) @@ -223,6 +237,10 @@ class WebRtc { void ProcessRestartTachyonReceiveMessages(const std::string& service_id) ABSL_LOCKS_EXCLUDED(mutex_); + // Runs on |single_thread_executor_|. + void RestartTachyonReceiveMessages(const std::string& service_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void OffloadFromThread(Runnable runnable); Mutex mutex_; diff --git a/cpp/core/internal/mediums/webrtc_test.cc b/cpp/core/internal/mediums/webrtc_test.cc index 535a11b9..01e49045 100644 --- a/cpp/core/internal/mediums/webrtc_test.cc +++ b/cpp/core/internal/mediums/webrtc_test.cc @@ -158,9 +158,9 @@ TEST_F(WebRtcTest, StartAcceptingConnectionTwice) { EXPECT_FALSE(webrtc.IsAcceptingConnections(std::string{})); } -// Tests the flow when the device tries to connect but the data channel times -// out. -TEST_F(WebRtcTest, Connect_DataChannelTimeOut) { +// Tests the flow when the device tries to connect but there is no peer +// accepting connections at the given peer ID. +TEST_F(WebRtcTest, Connect_NoPeer) { WebRtc webrtc; PeerId peer_id("peer_id"); const std::string service_id("NearbySharing"); @@ -353,6 +353,39 @@ TEST_F(WebRtcTest, Connect_NullPeerConnection) { EXPECT_FALSE(wrapper.IsValid()); } +// Tests the flow when the device calls StartAcceptingConnections and the +// receive messages stream fails. +TEST_F(WebRtcTest, ContinueAcceptingConnectionsOnComplete) { + using MockAcceptedCallback = + testing::MockFunction; + testing::StrictMock mock_accepted_callback_; + + WebRtc webrtc; + PeerId self_id("peer_id"); + const std::string service_id("NearbySharing"); + LocationHint location_hint; + + ASSERT_TRUE(webrtc.IsAvailable()); + ASSERT_TRUE(webrtc.StartAcceptingConnections( + service_id, self_id, location_hint, + {mock_accepted_callback_.AsStdFunction()})); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); + + // Simulate a failure in receiving messages stream, WebRtc should restart + // accepting connections. + MediumEnvironment::Instance().SendWebRtcSignalingComplete(self_id.GetId(), + /*success=*/false); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); + + // And a "success" message should not cause accepting connections to stop. + MediumEnvironment::Instance().SendWebRtcSignalingComplete(self_id.GetId(), + /*success=*/true); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); + + webrtc.StopAcceptingConnections(service_id); + EXPECT_FALSE(webrtc.IsAcceptingConnections(service_id)); +} + } // namespace } // namespace mediums diff --git a/cpp/platform/api/webrtc.h b/cpp/platform/api/webrtc.h index 4ddfdb12..16fc30e5 100644 --- a/cpp/platform/api/webrtc.h +++ b/cpp/platform/api/webrtc.h @@ -30,13 +30,16 @@ namespace api { class WebRtcSignalingMessenger { public: using OnSignalingMessageCallback = std::function; + using OnSignalingCompleteCallback = std::function; virtual ~WebRtcSignalingMessenger() = default; virtual bool SendMessage(absl::string_view peer_id, const ByteArray& message) = 0; - virtual bool StartReceivingMessages(OnSignalingMessageCallback listener) = 0; + virtual bool StartReceivingMessages( + OnSignalingMessageCallback on_message_callback, + OnSignalingCompleteCallback on_complete_callback) = 0; virtual void StopReceivingMessages() = 0; }; diff --git a/cpp/platform/base/medium_environment.cc b/cpp/platform/base/medium_environment.cc index 497007cb..42733cd8 100644 --- a/cpp/platform/base/medium_environment.cc +++ b/cpp/platform/base/medium_environment.cc @@ -59,6 +59,8 @@ void MediumEnvironment::Reset() { bluetooth_adapters_.clear(); bluetooth_mediums_.clear(); ble_mediums_.clear(); + webrtc_signaling_message_callback_.clear(); + webrtc_signaling_complete_callback_.clear(); wifi_lan_mediums_.clear(); }); Sync(); @@ -439,11 +441,17 @@ void MediumEnvironment::CallBleAcceptedConnectionCallback( } void MediumEnvironment::RegisterWebRtcSignalingMessenger( - absl::string_view self_id, OnSignalingMessageCallback callback) { + absl::string_view self_id, OnSignalingMessageCallback message_callback, + OnSignalingCompleteCallback complete_callback) { if (!enabled_) return; RunOnMediumEnvironmentThread( - [this, self_id{std::string(self_id)}, callback{std::move(callback)}]() { - webrtc_signaling_callback_[self_id] = std::move(callback); + [this, self_id{std::string(self_id)}, + message_callback{std::move(message_callback)}, + complete_callback{std::move(complete_callback)}]() { + webrtc_signaling_message_callback_[self_id] = + std::move(message_callback); + webrtc_signaling_complete_callback_[self_id] = + std::move(complete_callback); NEARBY_LOG(INFO, "Registered signaling message callback for id = %s", self_id.c_str()); }); @@ -453,9 +461,11 @@ void MediumEnvironment::UnregisterWebRtcSignalingMessenger( absl::string_view self_id) { if (!enabled_) return; RunOnMediumEnvironmentThread([this, self_id{std::string(self_id)}]() { - auto item = webrtc_signaling_callback_.extract(self_id); - if (item.empty()) return; - NEARBY_LOG(INFO, "Unregistered signaling message callback for id = %s", + auto message_callback_item = + webrtc_signaling_message_callback_.extract(self_id); + auto complete_callback_item = + webrtc_signaling_complete_callback_.extract(self_id); + NEARBY_LOG(INFO, "Unregistered signaling callbacks for id = %s", self_id.c_str()); }); } @@ -465,8 +475,8 @@ void MediumEnvironment::SendWebRtcSignalingMessage(absl::string_view peer_id, if (!enabled_) return; RunOnMediumEnvironmentThread( [this, peer_id{std::string(peer_id)}, message]() { - auto item = webrtc_signaling_callback_.find(peer_id); - if (item == webrtc_signaling_callback_.end()) { + auto item = webrtc_signaling_message_callback_.find(peer_id); + if (item == webrtc_signaling_message_callback_.end()) { NEARBY_LOG(WARNING, "No callback registered for peer id = %s", peer_id.c_str()); return; @@ -476,6 +486,22 @@ void MediumEnvironment::SendWebRtcSignalingMessage(absl::string_view peer_id, }); } +void MediumEnvironment::SendWebRtcSignalingComplete(absl::string_view peer_id, + bool success) { + if (!enabled_) return; + RunOnMediumEnvironmentThread( + [this, peer_id{std::string(peer_id)}, success]() { + auto item = webrtc_signaling_complete_callback_.find(peer_id); + if (item == webrtc_signaling_complete_callback_.end()) { + NEARBY_LOG(WARNING, "No callback registered for peer id = %s", + peer_id.c_str()); + return; + } + + item->second(success); + }); +} + void MediumEnvironment::SetUseValidPeerConnection( bool use_valid_peer_connection) { use_valid_peer_connection_ = use_valid_peer_connection; diff --git a/cpp/platform/base/medium_environment.h b/cpp/platform/base/medium_environment.h index e215bffc..bd7126dc 100644 --- a/cpp/platform/base/medium_environment.h +++ b/cpp/platform/base/medium_environment.h @@ -55,6 +55,8 @@ class MediumEnvironment { api::BleMedium::AcceptedConnectionCallback; using OnSignalingMessageCallback = api::WebRtcSignalingMessenger::OnSignalingMessageCallback; + using OnSignalingCompleteCallback = + api::WebRtcSignalingMessenger::OnSignalingCompleteCallback; using WifiLanDiscoveredServiceCallback = api::WifiLanMedium::DiscoveredServiceCallback; using WifiLanAcceptedConnectionCallback = @@ -128,9 +130,11 @@ class MediumEnvironment { const EnvironmentConfig& GetEnvironmentConfig(); - // Registers |callback| to receive messages sent to device with id |self_id|. - void RegisterWebRtcSignalingMessenger(absl::string_view self_id, - OnSignalingMessageCallback callback); + // Registers |message_callback| to receive messages sent to device with id + // |self_id|, and |complete_callback| to notify when signaling is complete. + void RegisterWebRtcSignalingMessenger( + absl::string_view self_id, OnSignalingMessageCallback message_callback, + OnSignalingCompleteCallback complete_callback); // Unregisters the callback listening to incoming messages for |self_id|. void UnregisterWebRtcSignalingMessenger(absl::string_view self_id); @@ -140,6 +144,9 @@ class MediumEnvironment { void SendWebRtcSignalingMessage(absl::string_view peer_id, const ByteArray& message); + // Simulates sending an "signaling complete" signal to the WebRTC medium. + void SendWebRtcSignalingComplete(absl::string_view peer_id, bool success); + // Used to set if WebRtcMedium should use a valid peer connection or nullptr // in tests. void SetUseValidPeerConnection(bool use_valid_peer_connection); @@ -302,7 +309,11 @@ class MediumEnvironment { // Maps peer id to callback for receiving signaling messages. absl::flat_hash_map - webrtc_signaling_callback_; + webrtc_signaling_message_callback_; + + // Maps peer id to callback for signaling complete events. + absl::flat_hash_map + webrtc_signaling_complete_callback_; absl::flat_hash_map wifi_lan_mediums_; diff --git a/cpp/platform/impl/g3/webrtc.cc b/cpp/platform/impl/g3/webrtc.cc index 858b23f6..cd5f467c 100644 --- a/cpp/platform/impl/g3/webrtc.cc +++ b/cpp/platform/impl/g3/webrtc.cc @@ -35,9 +35,11 @@ bool WebRtcSignalingMessenger::SendMessage(absl::string_view peer_id, } bool WebRtcSignalingMessenger::StartReceivingMessages( - OnSignalingMessageCallback listener) { + OnSignalingMessageCallback on_message_callback, + OnSignalingCompleteCallback on_complete_callback) { auto& env = MediumEnvironment::Instance(); - env.RegisterWebRtcSignalingMessenger(self_id_, listener); + env.RegisterWebRtcSignalingMessenger(self_id_, on_message_callback, + on_complete_callback); return true; } diff --git a/cpp/platform/impl/g3/webrtc.h b/cpp/platform/impl/g3/webrtc.h index 20f310c8..235fb649 100644 --- a/cpp/platform/impl/g3/webrtc.h +++ b/cpp/platform/impl/g3/webrtc.h @@ -29,6 +29,8 @@ class WebRtcSignalingMessenger : public api::WebRtcSignalingMessenger { public: using OnSignalingMessageCallback = api::WebRtcSignalingMessenger::OnSignalingMessageCallback; + using OnSignalingCompleteCallback = + api::WebRtcSignalingMessenger::OnSignalingCompleteCallback; explicit WebRtcSignalingMessenger( absl::string_view self_id, @@ -37,7 +39,9 @@ class WebRtcSignalingMessenger : public api::WebRtcSignalingMessenger { bool SendMessage(absl::string_view peer_id, const ByteArray& message) override; - bool StartReceivingMessages(OnSignalingMessageCallback listener) override; + bool StartReceivingMessages( + OnSignalingMessageCallback on_message_callback, + OnSignalingCompleteCallback on_complete_callback) override; void StopReceivingMessages() override; private: diff --git a/cpp/platform/public/webrtc.h b/cpp/platform/public/webrtc.h index 5d28d2e1..56b497ed 100644 --- a/cpp/platform/public/webrtc.h +++ b/cpp/platform/public/webrtc.h @@ -28,6 +28,8 @@ class WebRtcSignalingMessenger final { public: using OnSignalingMessageCallback = api::WebRtcSignalingMessenger::OnSignalingMessageCallback; + using OnSignalingCompleteCallback = + api::WebRtcSignalingMessenger::OnSignalingCompleteCallback; explicit WebRtcSignalingMessenger( std::unique_ptr messenger) @@ -40,8 +42,11 @@ class WebRtcSignalingMessenger final { return impl_->SendMessage(peer_id, message); } - bool StartReceivingMessages(OnSignalingMessageCallback listener) { - return impl_->StartReceivingMessages(listener); + bool StartReceivingMessages( + OnSignalingMessageCallback on_message_callback, + OnSignalingCompleteCallback on_complete_callback) { + return impl_->StartReceivingMessages(on_message_callback, + on_complete_callback); } void StopReceivingMessages() { impl_->StopReceivingMessages(); }