From 16cb7f70b3a8b8b04440b7b96534e11df964ca60 Mon Sep 17 00:00:00 2001 From: hai007 Date: Tue, 19 Dec 2023 14:26:19 -0800 Subject: [PATCH] Implement Auto-reconnect after disconnection [1] PiperOrigin-RevId: 592344338 --- Package.swift | 2 + connections/implementation/BUILD | 6 + .../implementation/bluetooth_bwu_handler.cc | 1 + .../implementation/bluetooth_bwu_test.cc | 126 +++ connections/implementation/client_proxy.cc | 44 +- connections/implementation/client_proxy.h | 19 + .../flags/nearby_connections_feature_flags.h | 4 + connections/implementation/offline_frames.cc | 28 + connections/implementation/offline_frames.h | 2 + .../implementation/offline_frames_test.cc | 37 + .../implementation/p2p_cluster_pcp_handler.cc | 1 + .../implementation/reconnect_manager.cc | 864 ++++++++++++++++++ .../implementation/reconnect_manager.h | 237 +++++ .../implementation/reconnect_manager_test.cc | 149 +++ .../implementation/service_id_constants.h | 33 + connections/implementation/simulation_user.h | 21 +- internal/platform/feature_flags.h | 8 + internal/platform/implementation/BUILD | 1 + 18 files changed, 1576 insertions(+), 7 deletions(-) create mode 100644 connections/implementation/bluetooth_bwu_test.cc create mode 100644 connections/implementation/reconnect_manager.cc create mode 100644 connections/implementation/reconnect_manager.h create mode 100644 connections/implementation/reconnect_manager_test.cc diff --git a/Package.swift b/Package.swift index 713cef3e..394ea817 100644 --- a/Package.swift +++ b/Package.swift @@ -429,6 +429,7 @@ let package = Package( "connections/implementation/payload_manager_test.cc", "connections/implementation/offline_frames_validator_test.cc", "connections/implementation/service_controller_router_test.cc", + "connections/implementation/bluetooth_bwu_test.cc", "connections/implementation/wifi_direct_bwu_test.cc", "connections/implementation/wifi_hotspot_test.cc", "connections/implementation/analytics/analytics_recorder_test.cc", @@ -461,6 +462,7 @@ let package = Package( "connections/implementation/pcp_manager_test.cc", "connections/implementation/ble_advertisement_test.cc", "connections/implementation/base_endpoint_channel_test.cc", + "connections/implementation/reconnect_manager_test.cc", "connections/v3/connections_device_test.cc", "connections/v3/connections_device_provider_test.cc", "connections/implementation/connections_authentication_transport_test.cc", diff --git a/connections/implementation/BUILD b/connections/implementation/BUILD index 4a42d633..2564fbf8 100644 --- a/connections/implementation/BUILD +++ b/connections/implementation/BUILD @@ -42,6 +42,7 @@ cc_library( "p2p_star_pcp_handler.cc", "payload_manager.cc", "pcp_manager.cc", + "reconnect_manager.cc", "service_controller_router.cc", "webrtc_bwu_handler.cc", "webrtc_bwu_handler_stub.cc", @@ -85,6 +86,7 @@ cc_library( "pcp.h", "pcp_handler.h", "pcp_manager.h", + "reconnect_manager.h", "service_controller.h", "service_controller_router.h", "service_id_constants.h", @@ -130,6 +132,7 @@ cc_library( "//internal/platform:util", "//internal/platform/implementation:comm", "//internal/platform/implementation:platform", + "//internal/platform/implementation:types", "//internal/proto/analytics:connections_log_cc_proto", "//proto:connections_enums_cc_proto", "@com_google_absl//absl/base:core_headers", @@ -138,6 +141,7 @@ cc_library( "@com_google_absl//absl/container:flat_hash_set", "@com_google_absl//absl/functional:any_invocable", "@com_google_absl//absl/functional:bind_front", + "@com_google_absl//absl/log:check", "@com_google_absl//absl/memory", "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:str_format", @@ -204,6 +208,7 @@ cc_test( "base_endpoint_channel_test.cc", "base_pcp_handler_test.cc", "ble_advertisement_test.cc", + "bluetooth_bwu_test.cc", "bluetooth_device_name_test.cc", "bwu_manager_test.cc", "client_proxy_test.cc", @@ -219,6 +224,7 @@ cc_test( "p2p_point_to_point_pcp_handler_test.cc", "payload_manager_test.cc", "pcp_manager_test.cc", + "reconnect_manager_test.cc", "service_controller_router_test.cc", "wifi_direct_bwu_test.cc", "wifi_hotspot_test.cc", diff --git a/connections/implementation/bluetooth_bwu_handler.cc b/connections/implementation/bluetooth_bwu_handler.cc index 21328239..7b71e610 100644 --- a/connections/implementation/bluetooth_bwu_handler.cc +++ b/connections/implementation/bluetooth_bwu_handler.cc @@ -93,6 +93,7 @@ BluetoothBwuHandler::CreateUpgradedEndpointChannel( return nullptr; } + client->SetBluetoothMacAddress(endpoint_id, mac_address); return channel; } diff --git a/connections/implementation/bluetooth_bwu_test.cc b/connections/implementation/bluetooth_bwu_test.cc new file mode 100644 index 00000000..6b0b9897 --- /dev/null +++ b/connections/implementation/bluetooth_bwu_test.cc @@ -0,0 +1,126 @@ +// Copyright 2022 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include + +#include "gtest/gtest.h" +#include "absl/time/time.h" +#include "connections/implementation/bwu_handler.h" +#include "connections/implementation/bluetooth_bwu_handler.h" +#include "connections/implementation/client_proxy.h" +#include "connections/implementation/endpoint_channel.h" +#include "connections/implementation/mediums/mediums.h" +#include "connections/implementation/offline_frames.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/count_down_latch.h" +#include "internal/platform/exception.h" +#include "internal/platform/feature_flags.h" +#include "internal/platform/logging.h" +#include "internal/platform/medium_environment.h" +#include "internal/platform/single_thread_executor.h" + +namespace nearby { +namespace connections { + +namespace { +using ::location::nearby::connections::OfflineFrame; +constexpr absl::Duration kWaitDuration = absl::Milliseconds(1000); +} // namespace + +class BluetoothBwuTest : public testing::Test { + protected: + BluetoothBwuTest() { env_.Start(); } + ~BluetoothBwuTest() override { env_.Stop(); } + + MediumEnvironment& env_{MediumEnvironment::Instance()}; +}; + +TEST_F(BluetoothBwuTest, CanCreateBwuHandler) { + ClientProxy client; + Mediums mediums; + + auto handler = std::make_unique(mediums, nullptr); + + handler->InitializeUpgradedMediumForEndpoint(&client, /*service_id=*/"B", + /*endpoint_id=*/"2"); + handler->RevertInitiatorState(); + SUCCEED(); + handler.reset(); +} + +TEST_F(BluetoothBwuTest, SoftAPBWUInit_STACreateEndpointChannel) { + CountDownLatch start_latch(1); + CountDownLatch accept_latch(1); + CountDownLatch end_latch(1); + + ClientProxy client_1, client_2; + Mediums mediums_1, mediums_2; + ExceptionOr upgrade_frame; + + auto handler_1 = std::make_unique( + mediums_1, [&](ClientProxy* client, + std::unique_ptr + mutable_connection) { + NEARBY_LOGS(WARNING) << "Server socket connection accept call back"; + accept_latch.CountDown(); + EXPECT_TRUE(end_latch.Await(kWaitDuration).result()); + }); + + // client_1 works as Bluetooth Server Device + SingleThreadExecutor server_executor; + server_executor.Execute([&]() { + ByteArray upgrade_path_available_frame = + handler_1->InitializeUpgradedMediumForEndpoint(&client_1, + /*service_id=*/"A", + /*endpoint_id=*/"1"); + EXPECT_FALSE(upgrade_path_available_frame.Empty()); + + upgrade_frame = parser::FromBytes(upgrade_path_available_frame); + start_latch.CountDown(); + }); + + // client_2 works as Bluetooth Client Device which will connect to client_1 + SingleThreadExecutor client_executor; + // Wait till client_1 started as Bluetooth and then connect to it + EXPECT_TRUE(start_latch.Await(kWaitDuration).result()); + std::unique_ptr handler_2 = + std::make_unique(mediums_2, nullptr); + + client_executor.Execute([&]() { + auto bwu_frame = + upgrade_frame.result().v1().bandwidth_upgrade_negotiation(); + + std::unique_ptr new_channel = + handler_2->CreateUpgradedEndpointChannel(&client_2, /*service_id=*/"A", + /*endpoint_id=*/"1", + bwu_frame.upgrade_path_info()); + if (!FeatureFlags::GetInstance().GetFlags().enable_cancellation_flag) { + EXPECT_TRUE(accept_latch.Await(kWaitDuration).result()); + EXPECT_EQ(new_channel->GetMedium(), + location::nearby::proto::connections::Medium::BLUETOOTH); + } else { + accept_latch.CountDown(); + EXPECT_EQ(new_channel, nullptr); + } + EXPECT_FALSE(mediums_2.GetBluetoothClassic().GetMacAddress().empty()); + handler_2->RevertResponderState(/*service_id=*/"A"); + end_latch.CountDown(); + }); + + EXPECT_TRUE(accept_latch.Await(kWaitDuration).result()); + EXPECT_TRUE(end_latch.Await(kWaitDuration).result()); +} + +} // namespace connections +} // namespace nearby diff --git a/connections/implementation/client_proxy.cc b/connections/implementation/client_proxy.cc index c86ae99f..3cd87094 100644 --- a/connections/implementation/client_proxy.cc +++ b/connections/implementation/client_proxy.cc @@ -74,12 +74,14 @@ ClientProxy::ClientProxy(::nearby::analytics::EventLogger* event_logger) supports_safe_to_disconnect_ = NearbyFlags::GetInstance().GetBoolFlag( config_package_nearby::nearby_connections_feature:: kEnableSafeToDisconnect); + support_auto_reconnect_ = NearbyFlags::GetInstance().GetBoolFlag( + config_package_nearby::nearby_connections_feature::kEnableAutoReconnect); local_safe_to_disconnect_version_ = NearbyFlags::GetInstance().GetInt64Flag( config_package_nearby::nearby_connections_feature:: kSafeToDisconnectVersion); NEARBY_LOGS(INFO) << "[safe-to-disconnect]: Local enabled: " << supports_safe_to_disconnect_ - << "; Version_: " << local_safe_to_disconnect_version_; + << "; Version: " << local_safe_to_disconnect_version_; } ClientProxy::~ClientProxy() { Reset(); } @@ -120,6 +122,18 @@ std::string ClientProxy::GetConnectionToken(const std::string& endpoint_id) { return {}; } +std::optional ClientProxy::GetBluetoothMacAddress( + const std::string& endpoint_id) { + auto item = bluetooth_mac_addresses_.find(endpoint_id); + if (item != bluetooth_mac_addresses_.end()) return item->second; + return std::nullopt; +} + +void ClientProxy::SetBluetoothMacAddress( + const std::string& endpoint_id, const std::string& bluetooth_mac_address) { + bluetooth_mac_addresses_[endpoint_id] = bluetooth_mac_address; +} + std::string ClientProxy::GenerateLocalEndpointId() { if (high_vis_mode_) { if (!local_high_vis_mode_cache_endpoint_id_.empty()) { @@ -611,6 +625,24 @@ std::int32_t ClientProxy::GetNumIncomingConnections() const { .size(); } +bool ClientProxy::IsIncomingConnection(const std::string& endpoint_id) const { + MutexLock lock(&mutex_); + const ConnectionPair* item = LookupConnection(endpoint_id); + if (item != nullptr && item->first.status == Connection::kConnected) { + return item->first.is_incoming; + } + return false; +} + +bool ClientProxy::IsOutgoingConnection(const std::string& endpoint_id) const { + MutexLock lock(&mutex_); + const ConnectionPair* item = LookupConnection(endpoint_id); + if (item != nullptr && item->first.status == Connection::kConnected) { + return !item->first.is_incoming; + } + return false; +} + bool ClientProxy::HasPendingConnectionToEndpoint( const std::string& endpoint_id) const { MutexLock lock(&mutex_); @@ -848,6 +880,15 @@ bool ClientProxy::IsSafeToDisconnectEnabled(absl::string_view endpoint_id) { .min_nc_version_supports_safe_to_disconnect); } +bool ClientProxy::IsAutoReconnectEnabled(absl::string_view endpoint_id) { + return IsSupportAutoReconnect() && + GetRemoteSafeToDisconnectVersion(endpoint_id).has_value() && + (GetRemoteSafeToDisconnectVersion(endpoint_id) >= + FeatureFlags::GetInstance() + .GetFlags() + .min_nc_version_supports_auto_reconnect); +} + bool ClientProxy::IsPayloadReceivedAckEnabled(absl::string_view endpoint_id) { return IsSupportSafeToDisconnect() && GetRemoteSafeToDisconnectVersion(endpoint_id).has_value() && @@ -929,6 +970,7 @@ void ClientProxy::RemoveAllEndpoints() { // just remove without notifying. connections_.clear(); cancellation_flags_.clear(); + bluetooth_mac_addresses_.clear(); OnSessionComplete(); } diff --git a/connections/implementation/client_proxy.h b/connections/implementation/client_proxy.h index ff6a8c4d..00067c0a 100644 --- a/connections/implementation/client_proxy.h +++ b/connections/implementation/client_proxy.h @@ -23,6 +23,7 @@ #include #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" #include "connections/advertising_options.h" #include "connections/discovery_options.h" #include "connections/implementation/analytics/analytics_recorder.h" @@ -74,6 +75,10 @@ class ClientProxy final { } std::string GetConnectionToken(const std::string& endpoint_id); + std::optional GetBluetoothMacAddress( + const std::string& endpoint_id); + void SetBluetoothMacAddress(const std::string& endpoint_id, + const std::string& bluetooth_mac_address); const NearbyDevice* GetLocalDevice(); NearbyDeviceProvider* GetLocalDeviceProvider() { if (external_device_provider_ != nullptr) { @@ -188,6 +193,10 @@ class ClientProxy final { std::int32_t GetNumOutgoingConnections() const; // Returns the number of endpoints that are connected and incoming. std::int32_t GetNumIncomingConnections() const; + // Returns true if endpoint is incoming connection. + bool IsIncomingConnection(const std::string& endpoint_id) const; + // Returns true if endpoint is outgoing connection. + bool IsOutgoingConnection(const std::string& endpoint_id) const; // If true, then we're in the process of approving (or rejecting) a // connection. No payloads should be sent until isConnectedToEndpoint() // returns true. @@ -270,6 +279,11 @@ class ClientProxy final { const bool& IsSupportSafeToDisconnect() const { return supports_safe_to_disconnect_; } + + bool IsSupportAutoReconnect() const { + return support_auto_reconnect_; + } + const std::int32_t& GetLocalSafeToDisconnectVersion() const { return local_safe_to_disconnect_version_; } @@ -279,6 +293,7 @@ class ClientProxy final { absl::string_view endpoint_id, const std::int32_t& safe_to_disconnect_version); bool IsSafeToDisconnectEnabled(absl::string_view endpoint_id); + bool IsAutoReconnectEnabled(absl::string_view endpoint_id); bool IsPayloadReceivedAckEnabled(absl::string_view endpoint_id); private: @@ -415,6 +430,9 @@ class ClientProxy final { // Maps endpoint_id to endpoint connection state. absl::flat_hash_map connections_; + // Maps endpoint_id to Bluetooth Mac Addresses. + absl::flat_hash_map bluetooth_mac_addresses_; + // A cache of endpoint ids that we've already notified the discoverer of. We // check this cache before calling onEndpointFound() so that we don't notify // the client multiple times for the same endpoint. This would otherwise @@ -443,6 +461,7 @@ class ClientProxy final { // For Nearby Connections' own device provider. std::unique_ptr connections_device_provider_; bool supports_safe_to_disconnect_; + bool support_auto_reconnect_; std::int32_t local_safe_to_disconnect_version_; }; diff --git a/connections/implementation/flags/nearby_connections_feature_flags.h b/connections/implementation/flags/nearby_connections_feature_flags.h index ecab2dd4..6990b431 100644 --- a/connections/implementation/flags/nearby_connections_feature_flags.h +++ b/connections/implementation/flags/nearby_connections_feature_flags.h @@ -49,6 +49,10 @@ constexpr auto kEnablePayloadManagerToSkipChunkUpdate = constexpr auto kEnableSafeToDisconnect = flags::Flag(kConfigPackage, "45425789", false); +// Enable/Disable auto_reconnect feature. +constexpr auto kEnableAutoReconnect = + flags::Flag(kConfigPackage, "45427690", false); + // When true, allows to enable payload-received-ack protocol. constexpr auto kEnablePayloadReceivedAck = flags::Flag(kConfigPackage, "45425840", false); diff --git a/connections/implementation/offline_frames.cc b/connections/implementation/offline_frames.cc index 5f5f86d9..ca468be2 100644 --- a/connections/implementation/offline_frames.cc +++ b/connections/implementation/offline_frames.cc @@ -23,6 +23,7 @@ #include "connections/implementation/flags/nearby_connections_feature_flags.h" #include "connections/implementation/offline_frames_validator.h" #include "connections/implementation/proto/offline_wire_formats.pb.h" +#include "connections/medium_selector.h" #include "connections/status.h" #include "internal/flags/nearby_flags.h" #include "internal/platform/byte_array.h" @@ -43,6 +44,7 @@ using ::location::nearby::connections::OfflineFrame; using ::location::nearby::connections::OsInfo; using ::location::nearby::connections::PayloadTransferFrame; using ::location::nearby::connections::V1Frame; +using ::location::nearby::connections::AutoReconnectFrame; ByteArray ToBytes(OfflineFrame&& frame) { ByteArray bytes(frame.ByteSizeLong()); @@ -469,6 +471,32 @@ ByteArray ForDisconnection(bool request_safe_to_disconnect, return ToBytes(std::move(frame)); } +ByteArray ForAutoReconnectIntroduction(const std::string& endpoint_id) { + OfflineFrame frame; + + frame.set_version(OfflineFrame::V1); + auto* v1_frame = frame.mutable_v1(); + v1_frame->set_type(V1Frame::AUTO_RECONNECT); + auto* auto_reconnect = v1_frame->mutable_auto_reconnect(); + auto_reconnect->set_endpoint_id(endpoint_id); + auto_reconnect->set_event_type(AutoReconnectFrame::CLIENT_INTRODUCTION); + + return ToBytes(std::move(frame)); +} + +ByteArray ForAutoReconnectIntroductionAck(const std::string& endpoint_id) { + OfflineFrame frame; + + frame.set_version(OfflineFrame::V1); + auto* v1_frame = frame.mutable_v1(); + v1_frame->set_type(V1Frame::AUTO_RECONNECT); + auto* auto_reconnect = v1_frame->mutable_auto_reconnect(); + auto_reconnect->set_endpoint_id(endpoint_id); + auto_reconnect->set_event_type(AutoReconnectFrame::CLIENT_INTRODUCTION_ACK); + + return ToBytes(std::move(frame)); +} + UpgradePathInfo::Medium MediumToUpgradePathInfoMedium(Medium medium) { switch (medium) { case Medium::MDNS: diff --git a/connections/implementation/offline_frames.h b/connections/implementation/offline_frames.h index 9d1a7d6f..c34fcc85 100644 --- a/connections/implementation/offline_frames.h +++ b/connections/implementation/offline_frames.h @@ -101,6 +101,8 @@ ByteArray ForBwuSafeToClose(); ByteArray ForKeepAlive(); ByteArray ForDisconnection(bool request_safe_to_disconnect, bool ack_safe_to_disconnect); +ByteArray ForAutoReconnectIntroduction(const std::string& endpoint_id); +ByteArray ForAutoReconnectIntroductionAck(const std::string& endpoint_id); UpgradePathInfo::Medium MediumToUpgradePathInfoMedium(Medium medium); Medium UpgradePathInfoMediumToMedium(UpgradePathInfo::Medium medium); diff --git a/connections/implementation/offline_frames_test.cc b/connections/implementation/offline_frames_test.cc index fea92f95..f70723cb 100644 --- a/connections/implementation/offline_frames_test.cc +++ b/connections/implementation/offline_frames_test.cc @@ -559,6 +559,43 @@ TEST(OfflineFramesTest, CanGenerateDisconnection) { EXPECT_THAT(message, EqualsProto(kExpected)); } +TEST(OfflineFramesTest, CanGenerateAutoReconnectIntroduction) { + constexpr absl::string_view kExpected = + R"pb( + version: V1 + v1: < + type: AUTO_RECONNECT + auto_reconnect: < + event_type: CLIENT_INTRODUCTION + endpoint_id: "ABC" + > + >)pb"; + ByteArray bytes = ForAutoReconnectIntroduction(std::string(kEndpointId)); + auto response = FromBytes(bytes); + ASSERT_TRUE(response.ok()); + OfflineFrame message = response.result(); + EXPECT_THAT(message, EqualsProto(kExpected)); +} + +TEST(OfflineFramesTest, CanGenerateAutoReconnectIntroductionAck) { + constexpr absl::string_view kExpected = + R"pb( + version: V1 + v1: < + type: AUTO_RECONNECT + auto_reconnect: < + event_type: CLIENT_INTRODUCTION_ACK + endpoint_id: "ABC" + > + >)pb"; + ByteArray bytes = ForAutoReconnectIntroductionAck(std::string(kEndpointId)); + auto response = FromBytes(bytes); + ASSERT_TRUE(response.ok()); + OfflineFrame message = response.result(); + EXPECT_THAT(message, EqualsProto(kExpected)); +} + + } // namespace } // namespace parser } // namespace connections diff --git a/connections/implementation/p2p_cluster_pcp_handler.cc b/connections/implementation/p2p_cluster_pcp_handler.cc index c06c6340..eefeed08 100644 --- a/connections/implementation/p2p_cluster_pcp_handler.cc +++ b/connections/implementation/p2p_cluster_pcp_handler.cc @@ -1641,6 +1641,7 @@ BasePcpHandler::ConnectImplResult P2pClusterPcpHandler::BluetoothConnectImpl( NEARBY_LOGS(VERBOSE) << "Client" << client->GetClientId() << " created Bluetooth endpoint channel to endpoint(id=" << endpoint->endpoint_id << ")."; + client->SetBluetoothMacAddress(endpoint->endpoint_id, device.GetMacAddress()); return BasePcpHandler::ConnectImplResult{ .medium = Medium::BLUETOOTH, .status = {Status::kSuccess}, diff --git a/connections/implementation/reconnect_manager.cc b/connections/implementation/reconnect_manager.cc new file mode 100644 index 00000000..f5e6b6cc --- /dev/null +++ b/connections/implementation/reconnect_manager.cc @@ -0,0 +1,864 @@ +// Copyright 2023 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "connections/implementation/reconnect_manager.h" + +#include +#include +#include +#include + +#include "securegcm/ukey2_handshake.h" +#include "absl/functional/any_invocable.h" +#include "absl/functional/bind_front.h" +#include "internal/platform/mutex.h" +#include "internal/platform/mutex_lock.h" +#ifndef NEARBY_CHROMIUM +#ifndef NEARBY_SWIFTPM +#include "absl/log/check.h" // nogncheck +#endif +#endif +#include "absl/strings/str_cat.h" +#include "absl/strings/string_view.h" +#include "connections/implementation/bluetooth_endpoint_channel.h" +#include "connections/implementation/client_proxy.h" +#include "connections/implementation/encryption_runner.h" +#include "connections/implementation/endpoint_channel.h" +#include "connections/implementation/endpoint_channel_manager.h" +#include "connections/implementation/mediums/mediums.h" +#include "connections/implementation/offline_frames.h" +#include "connections/implementation/service_id_constants.h" +#include "internal/platform/bluetooth_adapter.h" +#include "internal/platform/bluetooth_classic.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/cancelable_alarm.h" +#include "internal/platform/cancellation_flag_listener.h" +#include "internal/platform/count_down_latch.h" +#include "internal/platform/exception.h" +#include "internal/platform/feature_flags.h" +#include "internal/platform/implementation/system_clock.h" +#include "internal/platform/logging.h" +#include "proto/connections_enums.pb.h" + +namespace nearby { +namespace connections { +constexpr absl::string_view TAG = "[ReconnectManager]"; + +ReconnectManager::ReconnectManager(Mediums& mediums, + EndpointChannelManager& channel_manager) + : mediums_(&mediums), channel_manager_(&channel_manager) {} + +ReconnectManager::~ReconnectManager() { Shutdown(); } + +bool ReconnectManager::AutoReconnect( + ClientProxy* client, const std::string& endpoint_id, + AutoReconnectCallback& callback, + bool send_disconnection_notification, + DisconnectionReason disconnection_reason) { + if (!client->IsAutoReconnectEnabled(endpoint_id)) { + return false; + } + + if (resumed_endpoints_.contains(endpoint_id)) { + NEARBY_LOGS(INFO) << TAG << "AutoReconnect is not needed for endpoint_id = " + << endpoint_id + << ", since it's just reconnected successfully."; + return true; + } + + auto endpoint_channel = channel_manager_->GetChannelForEndpoint(endpoint_id); + if (endpoint_channel == nullptr) { + NEARBY_LOGS(INFO) + << TAG << " endpoint_channel shouldn't be null for endpoint_id = " + << endpoint_id; + return false; + } + Medium medium = endpoint_channel->GetMedium(); + + bool is_incoming = client->IsIncomingConnection(endpoint_id); + if (is_incoming == client->IsOutgoingConnection(endpoint_id)) { + NEARBY_LOGS(INFO) + << TAG << " autoReconnect failed for medium: " + << location::nearby::proto::connections::Medium_Name(medium) + << " because there is no existing incoming/outgoing connection, " + "is_incoming_connection = " + << is_incoming << ", is_outgoing_connection = " + << client->IsOutgoingConnection(endpoint_id); + return false; + } + std::string reconnect_service_id = + WrapInitiatorReconnectServiceId(endpoint_channel->GetServiceId()); + endpoint_id_metadata_map_.emplace( + endpoint_id, + ReconnectMetadata(is_incoming, std::move(callback), + send_disconnection_notification, disconnection_reason, + reconnect_service_id)); + NEARBY_LOGS(INFO) << TAG << "add a new endpoint_id " << endpoint_id + << " into metadata_by_service_id_map."; + + if (Start(is_incoming, client, endpoint_id, reconnect_service_id, medium)) { + resumed_endpoints_.emplace(endpoint_id); + + auto time_out = FeatureFlags::GetInstance() + .GetFlags() + .auto_reconnect_skip_duplicated_endpoint_duration; + std::make_unique( + absl::StrCat("RemoveSuccessfulResumedEndpointId for ", endpoint_id), + [this, endpoint_id, time_out]() { + NEARBY_LOGS(INFO) + << TAG << "Timeout after " << time_out + << "ms. RemoveSuccessfulResumedEndpointId for " << endpoint_id; + resumed_endpoints_.erase(endpoint_id); + }, + time_out, &alarm_executor_); + + return true; + } + + ClearReconnectData(client, reconnect_service_id, is_incoming); + return false; +} + +bool ReconnectManager::Start(bool is_incoming, ClientProxy* client, + const std::string& endpoint_id, + const std::string& reconnect_service_id, + Medium medium) { + auto retry_delay_millis = + FeatureFlags::GetInstance().GetFlags().auto_reconnect_retry_delay_millis; + auto reconnect_retry_num = + FeatureFlags::GetInstance().GetFlags().auto_reconnect_retry_attempts; + NEARBY_LOGS(INFO) << TAG << " " << (is_incoming ? "rehost" : "reconnect") + << " for medium: " + << location::nearby::proto::connections::Medium_Name( + medium) + << " for endpoint_id " << endpoint_id << " started..."; + bool final_result = false; + CountDownLatch latch(1); + reconnect_executor_.Execute( + "reconnect-start", + [this, &final_result, is_incoming, client, &endpoint_id, + &reconnect_service_id, retry_delay_millis, reconnect_retry_num, medium, + &latch]() mutable { + for (int i = 0; i < reconnect_retry_num; ++i) { + if (client->GetCancellationFlag(endpoint_id)->Cancelled()) { + NEARBY_LOGS(INFO) + << TAG << " Stop retry, Endpoint connection is cancelled"; + break; + } + if (RunOnce(is_incoming, client, endpoint_id, reconnect_service_id, + medium)) { + final_result = true; + break; + } + SystemClock::Sleep(retry_delay_millis); + } + NEARBY_LOGS(INFO) << "Reconnect " + << (final_result ? "succeeded" : "failed"); + latch.CountDown(); + }); + latch.Await(); + return final_result; +} + +bool ReconnectManager::RunOnce(bool is_incoming, ClientProxy* client, + const std::string& endpoint_id, + const std::string& reconnect_service_id, + Medium medium) { + bool result = false; + switch (medium) { + case Medium::BLUETOOTH: { + BluetoothImpl bluetooth_impl(client, endpoint_id, reconnect_service_id, + is_incoming, medium, mediums_, + channel_manager_, *this); + result = bluetooth_impl.Run(); + } break; + + default: + NEARBY_LOGS(INFO) << "AutoReconnect not implemented yet for " + << location::nearby::proto::connections::Medium_Name( + medium); + + break; + } + return result; +} + +void ReconnectManager::ClearReconnectData( + ClientProxy* client, const std::string& reconnect_service_id, + bool is_incoming) { + for (auto& item : endpoint_id_metadata_map_) { + if (item.second.reconnect_service_id == reconnect_service_id && + item.second.is_incoming == is_incoming) { + if (item.second.reconnect_cb.on_reconnect_failure_cb) { + item.second.reconnect_cb.on_reconnect_failure_cb( + client, item.first, item.second.send_disconnection_notification, + item.second.disconnection_reason); + } + } + NEARBY_LOGS(INFO) << TAG << "erase endpoint_id " << item.first; + endpoint_id_metadata_map_.erase(item.first); + } +} + +void ReconnectManager::Shutdown() { + NEARBY_LOGS(INFO) << TAG << "Initiating shutdown of ReconnectManager."; + { + MutexLock lock(&mutex_); + listen_timeout_alarm_by_service_id_.clear(); + } + new_endpoint_channels_.clear(); + endpoint_id_metadata_map_.clear(); + resumed_endpoints_.clear(); + + alarm_executor_.Shutdown(); + reconnect_executor_.Shutdown(); + encryption_cb_executor_.Shutdown(); + incoming_connection_cb_executor_.Shutdown(); + NEARBY_LOGS(INFO) << TAG << "ReconnectManager has shut down."; +} + +bool ReconnectManager::BaseMediumImpl::Run() { + if (!IsMediumRadioOn()) { + NEARBY_LOGS(INFO) << TAG + << location::nearby::proto::connections::Medium_Name( + medium_) + << " radio is turned off, try later"; + return false; + } + + if (client_->IsConnectedToEndpoint(endpoint_id_)) { + NEARBY_LOGS(INFO) << TAG + << "ReconnectBluetooth is not needed since it's already " + "connected to the RemoteDevice: "; + return true; + } + + auto previou_channel = channel_manager_->GetChannelForEndpoint(endpoint_id_); + if (previou_channel == nullptr) { + NEARBY_LOGS(INFO) + << TAG + << "ReconnectionManager didn't find a previous EndpointChannel " + "for " + << endpoint_id_ << " in this run, stop Reconnection!"; + return false; + } + previou_channel->Close( + DisconnectionReason::PREV_CHANNEL_DISCONNECTION_IN_RECONNECT); + + return is_incoming_ ? RehostForIncomingConnections() + : ReconnectToRemoteDevice(); +} + +bool ReconnectManager::BaseMediumImpl::RehostForIncomingConnections() { + auto time_out = + FeatureFlags::GetInstance().GetFlags().auto_reconnect_timeout_millis; + auto cancellation_flag = client_->GetCancellationFlag(endpoint_id_); + if (!IsListeningForIncomingConnections()) { + NEARBY_LOGS(INFO) << "Start rehosting for: " << reconnect_service_id_; + if (!StartListeningForIncomingConnections()) { + NEARBY_LOGS(ERROR) + << TAG + << "Rehost failed since " + "StartListeningForIncomingConnections return false."; + return false; + } + { + MutexLock lock(&reconnect_manager_.mutex_); + reconnect_manager_ + .listen_timeout_alarm_by_service_id_[reconnect_service_id_] = + std::make_unique( + absl::StrCat("Rehost listen timeout for ", reconnect_service_id_), + [this, time_out]() { + NEARBY_LOGS(INFO) + << "Timeout after " << time_out + << "ms. Stop listening for incoming " + "Connections for serviceId " + << reconnect_service_id_ << " for rehost, initiated by " + << endpoint_id_ + << ", unregister all still not connected endpointIds."; + StopListeningIfAllConnected( + reconnect_service_id_, + [this]() { StopListeningForIncomingConnections(); }, + /* forceStop= */ true); + }, + time_out, &reconnect_manager_.alarm_executor_); + } + } else { + NEARBY_LOGS(INFO) << "Rehosting is not needed since it's already " + "rehosts for: " + << reconnect_service_id_; + } + + if (cancellation_flag == nullptr) { + return true; + } + if (cancellation_flag->Cancelled()) { + StopListeningIfAllConnected( + reconnect_service_id_, + [this]() { StopListeningForIncomingConnections(); }, + /* forceStop= */ false); + return false; + } + auto cancellation_listener = + std::make_unique( + cancellation_flag, [this]() { + NEARBY_LOGS(INFO) << "Calling CancellationFlagListener."; + ProcessFailedReconnection(endpoint_id_, [this]() { + StopListeningForIncomingConnections(); + }); + }); + + std::make_unique( + absl::StrCat(TAG, " unregisterOnCancelListener"), + [cancellation_listener = std::move(cancellation_listener)]() mutable { + // clean up the listener after auto reconnect is done. + cancellation_listener.reset(); + }, + time_out, &reconnect_manager_.alarm_executor_); + return true; +} + +bool ReconnectManager::BaseMediumImpl::ReconnectToRemoteDevice() { + if (!ConnectOverMedium()) { + NEARBY_LOGS(INFO) << TAG << "Connect over medium " + << location::nearby::proto::connections::Medium_Name( + medium_) + << " failed."; + return false; + } + NEARBY_LOGS(INFO) << TAG << "Write CLIENT_INTRODUCTION frame"; + Exception write_exception = reconnect_channel_->Write( + parser::ForAutoReconnectIntroduction(endpoint_id_)); + if (!write_exception.Ok()) { + NEARBY_LOGS(ERROR) + << TAG << "Failed to write forAutoReconnectClientIntroductionEvent."; + QuietlyCloseChannelAndSocket(); + return false; + } + if (!ReadClientIntroductionAckFrame(reconnect_channel_.get())) { + NEARBY_LOGS(ERROR) << TAG << "Failed to read ClientIntroductionAck frame."; + QuietlyCloseChannelAndSocket(); + return false; + } + if (ReplaceChannelForEndpoint(client_, endpoint_id_, + std::move(reconnect_channel_), + SupportEncryptionDisabled(), nullptr)) { + NEARBY_LOGS(INFO) << TAG + << " successfully rebuild the outgoing connection with " + << location::nearby::proto::connections::Medium_Name( + medium_) + << " for the endpointId:" << endpoint_id_; + return true; + } + NEARBY_LOGS(INFO) + << TAG << " ReplaceChannelForEndpoint for the outgoing connection with " + << location::nearby::proto::connections::Medium_Name(medium_) + << " for the endpointId:" << endpoint_id_ << " failed. Please retry"; + return false; +} + +void ReconnectManager::BaseMediumImpl::OnIncomingConnection( + const std::string& reconnect_service_id) { + NEARBY_LOGS(INFO) << TAG << "Received reconnection successfully"; + reconnect_manager_.incoming_connection_cb_executor_.Execute( + "OnIncomingConnection", [this]() { + auto incoming_endpoin_id = + ReadClientIntroductionFrame(reconnect_channel_.get()); + if (incoming_endpoin_id.empty()) { + NEARBY_LOGS(ERROR) << TAG << "read ClientIntroductionFrame failed"; + QuietlyCloseChannelAndSocket(); + return; + } + NEARBY_LOGS(INFO) << TAG << "Write CLIENT_INTRODUCTION_ACK frame"; + Exception write_exception = reconnect_channel_->Write( + parser::ForAutoReconnectIntroductionAck(endpoint_id_)); + if (!write_exception.Ok()) { + NEARBY_LOGS(ERROR) + << TAG + << "Failed to write forAutoReconnectClientIntroductionAckEvent."; + QuietlyCloseChannelAndSocket(); + return; + } + if (ReplaceChannelForEndpoint( + client_, endpoint_id_, std::move(reconnect_channel_), + SupportEncryptionDisabled(), + [this]() { StopListeningForIncomingConnections(); })) { + NEARBY_LOGS(INFO) + << TAG << " successfully rebuild the incoming connection with " + << location::nearby::proto::connections::Medium_Name(medium_) + << " for the endpointId:" << endpoint_id_; + return; + } + QuietlyCloseChannelAndSocket(); + NEARBY_LOGS(INFO) + << TAG + << " ReplaceChannelForEndpoint for the incoming connection with " + << location::nearby::proto::connections::Medium_Name(medium_) + << " for the endpointId:" << endpoint_id_ + << " failed. Please retry"; + return; + }); +} + +std::string ReconnectManager::BaseMediumImpl::ReadClientIntroductionFrame( + EndpointChannel* endpoint_channel) { + NEARBY_LOGS(INFO) << TAG << "Read CLIENT_INTRODUCTION frame"; + + auto timeout = FeatureFlags::GetInstance() + .GetFlags() + .safe_to_disconnect_auto_resume_timeout_millis; + CancelableAlarm timeout_alarm( + "ReconnectManager::ReadClientIntroductionFrame", + [timeout, endpoint_channel]() { + NEARBY_LOGS(ERROR) << "In ReconnectManager, failed to read the " + "ClientIntroductionFrame after " + << timeout + << ". Timing out and closing EndpointChannel " + << endpoint_channel->GetType(); + endpoint_channel->Close(); + }, + timeout, &reconnect_manager_.alarm_executor_); + + auto data = endpoint_channel->Read(); + timeout_alarm.Cancel(); + if (!data.ok()) { + NEARBY_LOGS(ERROR) + << "Data read fail when expecting a ClientIntroductionFrame from " + "EndpointChannel " + << endpoint_channel->GetType(); + return {}; + } + auto transfer(parser::FromBytes(data.result())); + if (!transfer.ok()) { + NEARBY_LOGS(ERROR) << "Attempted to read a ClientIntroductionFrame from " + "EndpointChannel " + << endpoint_channel->GetType() + << ", but was unable to obtain any OfflineFrame."; + return {}; + } + OfflineFrame frame = transfer.result(); + if (!frame.has_v1() || !frame.v1().has_auto_reconnect()) { + NEARBY_LOGS(ERROR) << "In ReadClientIntroductionFrame(), eExpected a " + "AUTO_RECONNECT v1 OfflineFrame but got a " + << parser::GetFrameType(frame) << " frame instead."; + return {}; + } + if (frame.v1().auto_reconnect().event_type() != + AutoReconnectFrame::CLIENT_INTRODUCTION) { + NEARBY_LOGS(ERROR) << "In ReadClientIntroductionFrame(), expected a " + "CLIENT_INTRODUCTION " + "v1 OfflineFrame but got a AUTO_RECONNECT frame " + "with eventType " + << frame.v1().auto_reconnect().event_type() + << " instead."; + return {}; + } + return frame.v1().auto_reconnect().endpoint_id(); +} + +bool ReconnectManager::BaseMediumImpl::ReadClientIntroductionAckFrame( + EndpointChannel* endpoint_channel) { + NEARBY_LOGS(INFO) << TAG << "Read CLIENT_INTRODUCTION_ACK frame"; + + auto timeout = FeatureFlags::GetInstance() + .GetFlags() + .safe_to_disconnect_auto_resume_timeout_millis; + CancelableAlarm timeout_alarm( + "ReconnectManager::ReadClientIntroductionAckFrame", + [timeout, endpoint_channel]() { + NEARBY_LOGS(ERROR) << "In ReconnectManager, failed to read the " + "ClientIntroductionAckFrame after " + << timeout + << ". Timing out and closing EndpointChannel " + << endpoint_channel->GetType(); + endpoint_channel->Close(); + }, + timeout, &reconnect_manager_.alarm_executor_); + + auto data = endpoint_channel->Read(); + timeout_alarm.Cancel(); + if (!data.ok()) return false; + auto transfer(parser::FromBytes(data.result())); + if (!transfer.ok()) { + NEARBY_LOGS(ERROR) << "Attempted to read a ClientIntroductionAckFrame from " + "EndpointChannel " + << endpoint_channel->GetType() + << ", but was unable to obtain any OfflineFrame."; + return false; + } + OfflineFrame frame = transfer.result(); + if (!frame.has_v1() || !frame.v1().has_auto_reconnect()) { + NEARBY_LOGS(ERROR) << "In ReadClientIntroductionAckFrame(), eExpected a " + "AUTO_RECONNECT v1 OfflineFrame but got a " + << parser::GetFrameType(frame) << " frame instead."; + return false; + } + if (frame.v1().auto_reconnect().event_type() != + AutoReconnectFrame::CLIENT_INTRODUCTION_ACK) { + NEARBY_LOGS(ERROR) << "In ReadClientIntroductionAckFrame(), expected a " + "CLIENT_INTRODUCTION_ACK " + "v1 OfflineFrame but got a AUTO_RECONNECT frame " + "with eventType " + << frame.v1().auto_reconnect().event_type() + << " instead."; + return false; + } + return true; +} + +bool ReconnectManager::BaseMediumImpl::ReplaceChannelForEndpoint( + ClientProxy* client, const std::string& endpoint_id, + std::unique_ptr new_channel, + bool support_encryption_disabled, + absl::AnyInvocable stop_listening_incoming_connection) { + auto& endpoint_id_metadata_map = reconnect_manager_.endpoint_id_metadata_map_; + auto reconnect_metadata = endpoint_id_metadata_map.find(endpoint_id); + if (reconnect_metadata == endpoint_id_metadata_map.end()) { + NEARBY_LOGS(ERROR) << TAG << "ReconnectMetadata is null for endpointId: " + << endpoint_id << " ,please retry!"; + return false; + } + + EndpointChannel* endpoint_channel = + reconnect_manager_.new_endpoint_channels_ + .emplace(endpoint_id, std::move(new_channel)) + .first->second.get(); + replace_channel_succeed_ = false; + wait_encryption_to_finish_ = std::make_unique(1); + if (reconnect_metadata->second.is_incoming) { + reconnect_manager_.encryption_runner_.StartServer( + client, endpoint_id, endpoint_channel, GetResultListener()); + } else { + reconnect_manager_.encryption_runner_.StartClient( + client, endpoint_id, endpoint_channel, GetResultListener()); + } + wait_encryption_to_finish_->Await( + FeatureFlags::GetInstance().GetFlags().auto_reconnect_timeout_millis); + + NEARBY_LOGS(INFO) << TAG + << "replace_channel_succeed_: " << replace_channel_succeed_ + << " for endpointId: " << endpoint_id; + + if (replace_channel_succeed_) { + ProcessSuccessfulReconnection( + endpoint_id, [this]() { StopListeningForIncomingConnections(); }); + client->GetAnalyticsRecorder().OnConnectionEstablished( + endpoint_id, endpoint_channel->GetMedium(), + client->GetConnectionToken(endpoint_id)); + } else { + ProcessFailedReconnection( + endpoint_id, [this]() { StopListeningForIncomingConnections(); }); + } + reconnect_manager_.new_endpoint_channels_.erase(endpoint_id); + return replace_channel_succeed_; +} + +EncryptionRunner::ResultListener +ReconnectManager::BaseMediumImpl::GetResultListener() { + return { + .on_success_cb = + [this](const std::string& endpoint_id, + std::unique_ptr ukey2, + const std::string& auth_token, + const ByteArray& raw_auth_token) { + reconnect_manager_.encryption_cb_executor_.Execute( + "encryption-success", + [this, endpoint_id, raw_ukey2 = ukey2.release(), auth_token, + raw_auth_token]() mutable { + OnEncryptionSuccessRunnable( + endpoint_id, + std::unique_ptr(raw_ukey2), + auth_token, raw_auth_token); + wait_encryption_to_finish_->CountDown(); + }); + }, + .on_failure_cb = + [this](const std::string& endpoint_id, EndpointChannel* channel) { + reconnect_manager_.encryption_cb_executor_.Execute( + "encryption-failure", [this, endpoint_id, channel]() mutable { + NEARBY_LOGS(ERROR) + << "Encryption failed for endpoint_id=" << endpoint_id + << " on medium=" + << location::nearby::proto::connections::Medium_Name( + channel->GetMedium()); + OnEncryptionFailureRunnable(endpoint_id, channel); + wait_encryption_to_finish_->CountDown(); + }); + }, + }; +} + +void ReconnectManager::BaseMediumImpl::OnEncryptionSuccessRunnable( + const std::string& endpoint_id, + std::unique_ptr ukey2, + const std::string& auth_token, const ByteArray& raw_auth_token) { + auto item = reconnect_manager_.new_endpoint_channels_.find(endpoint_id); + if (item == reconnect_manager_.new_endpoint_channels_.end()) { + NEARBY_LOGS(INFO) << "TAG" + << "OnEncryptionSuccess failed, new_endpoint_channel is " + "null for Endpoint:" + << endpoint_id; + return; + } + if (!ukey2) { + NEARBY_LOGS(INFO) + << "TAG" + << "OnEncryptionSuccess failed, ukey2 is null for Endpoint:" + << endpoint_id; + return; + } + + // After both parties accepted connection (presumably after verifying & + // matching security tokens), we are allowed to extract the shared key. + bool succeeded = ukey2->VerifyHandshake(); + CHECK(succeeded); // If this fails, it's a UKEY2 protocol bug. + auto context = ukey2->ToConnectionContext(); + CHECK(context); // there is no way how this can fail, if Verify succeeded. + // If it did, it's a UKEY2 protocol bug. + + if (!reconnect_manager_.channel_manager_->EncryptChannelForEndpoint( + endpoint_id, std::move(context))) { + NEARBY_LOGS(INFO) << "TAG" + << "new_endpoint_channel failed to update " + "EncryptionContext for Endpoint:" + << endpoint_id; + return; + } + auto previous_channel = + reconnect_manager_.channel_manager_->GetChannelForEndpoint(endpoint_id); + if (previous_channel == nullptr) { + NEARBY_LOGS(INFO) + << "TAG" + << "ReconnectionManager didn't find a previous EndpointChannel for " + << endpoint_id + << " when registering the new EndpointChannel, stop Reconnection!"; + item->second->Close(DisconnectionReason::UNFINISHED); + return; + } + reconnect_manager_.channel_manager_->ReplaceChannelForEndpoint( + client_, endpoint_id, std::move(item->second), + SupportEncryptionDisabled()); + replace_channel_succeed_ = true; +} + +void ReconnectManager::BaseMediumImpl::OnEncryptionFailureRunnable( + const std::string& endpoint_id, EndpointChannel* endpoint_channel) { + NEARBY_LOGS(INFO) + << "TAG" + << "new_endpoint_channel failed to use encryption for Endpoint:" + << endpoint_id; +} + +void ReconnectManager::BaseMediumImpl::ProcessSuccessfulReconnection( + const std::string& endpoint_id, + absl::AnyInvocable stop_listening_incoming_connection) { + auto& endpoint_id_metadata_map = reconnect_manager_.endpoint_id_metadata_map_; + auto reconnect_metadata = endpoint_id_metadata_map.find(endpoint_id); + if (reconnect_metadata == endpoint_id_metadata_map.end()) { + NEARBY_LOGS(ERROR) << TAG + << "when ProcessSuccessfulReconnection, endpoint_id: " + << endpoint_id + << " is already removed fromendpoint_id_metadata_map."; + return; + } + + auto medatdata = std::move(reconnect_metadata->second); + endpoint_id_metadata_map.erase(reconnect_metadata); + auto& callback = medatdata.reconnect_cb; + if (callback.on_reconnect_success_cb) { + callback.on_reconnect_success_cb(client_, endpoint_id); + } else { + NEARBY_LOGS(ERROR) << TAG + << "when ProcessSuccessfulReconnection, endpoint_id: " + << endpoint_id + << " callback.on_reconnect_success_cb is null"; + } + + if (medatdata.is_incoming && + stop_listening_incoming_connection) { + StopListeningIfAllConnected(medatdata.reconnect_service_id, + std::move(stop_listening_incoming_connection), + /* forceStop= */ false); + } +} +void ReconnectManager::BaseMediumImpl::ProcessFailedReconnection( + const std::string& endpoint_id, + absl::AnyInvocable stop_listening_incoming_connection) {} + +void ReconnectManager::BaseMediumImpl::StopListeningIfAllConnected( + const std::string& reconnect_service_id, + absl::AnyInvocable stop_listening_incoming_connection, + bool force_stop) { + if (!force_stop && HasPendingIncomingConnections(reconnect_service_id)) { + return; + } + CancelClearHostTimeoutAlarm(reconnect_service_id); + stop_listening_incoming_connection(); + ClearReconnectData(reconnect_service_id, /* is_incoming= */ true); + NEARBY_LOGS(INFO) << TAG + << " No more pending incoming connections, " + "stop_listening_incoming_connection for " + << reconnect_service_id << " before timeout."; +} + +bool ReconnectManager::BaseMediumImpl::HasPendingIncomingConnections( + const std::string& reconnect_service_id) { + for (auto& item : reconnect_manager_.endpoint_id_metadata_map_) { + if (item.second.reconnect_service_id == reconnect_service_id && + item.second.is_incoming) { + return true; + } + } + return false; +} + +void ReconnectManager::BaseMediumImpl:: + CancelClearHostTimeoutAlarm(const std::string& service_id) { + MutexLock lock(&reconnect_manager_.mutex_); + auto item = + reconnect_manager_.listen_timeout_alarm_by_service_id_.find(service_id); + if (item == reconnect_manager_.listen_timeout_alarm_by_service_id_.end()) + return; + + if (item->second->IsValid()) { + item->second->Cancel(); + item->second.reset(); + } + reconnect_manager_.listen_timeout_alarm_by_service_id_.erase(item); +} +void ReconnectManager::BaseMediumImpl:: + ClearReconnectData(const std::string& service_id, bool is_incoming) { + auto& metadata_map = reconnect_manager_.endpoint_id_metadata_map_; + for (auto item = metadata_map.begin(); item != metadata_map.end(); ) { + if (item->second.reconnect_service_id == service_id && is_incoming) { + auto& callback = item->second.reconnect_cb.on_reconnect_failure_cb; + if (callback) + callback(client_, item->first, + item->second.send_disconnection_notification, + item->second.disconnection_reason); + metadata_map.erase(item); + } else { + ++item; + } + } +} + +bool ReconnectManager::BluetoothImpl::IsMediumRadioOn() const { + return bluetooth_medium_.IsAvailable(); +} + +bool ReconnectManager::BluetoothImpl::IsListeningForIncomingConnections() + const { + return bluetooth_medium_.IsAcceptingConnections(reconnect_service_id_); +} + +bool ReconnectManager::BluetoothImpl::StartListeningForIncomingConnections() { + if (!bluetooth_medium_.StartAcceptingConnections( + reconnect_service_id_, + absl::bind_front( + &ReconnectManager::BluetoothImpl::OnIncomingBluetoothConnection, this, + client_))) { + NEARBY_LOGS(ERROR) + << "ReconnectManager::BluetoothImpl couldn't initiate the " + "BLUETOOTH reconnect for endpoint " + << endpoint_id_ + << " because it failed to start listening for " + "incoming Bluetooth connections."; + return false; + } + NEARBY_LOGS(INFO) << "ReconnectManager::BluetoothImpl successfully started " + "listening for incoming " + "reconnection on service_id=" + << reconnect_service_id_ << " for endpoint " + << endpoint_id_; + return true; +} + +void ReconnectManager::BluetoothImpl::OnIncomingBluetoothConnection( + ClientProxy* client, const std::string& upgrade_service_id, + BluetoothSocket socket) { + reconnect_channel_ = std::make_unique( + upgrade_service_id, /*channel_name=*/upgrade_service_id, socket); + if (reconnect_channel_ == nullptr) { + NEARBY_LOGS(ERROR) << TAG + << "Create new endpointChannel for incoming socket " + "failed, close the socket"; + + socket.Close(); + return; + } + bluetooth_socket_ = std::move(socket); + + NEARBY_LOGS(INFO) + << TAG << "Create new endpointChannel successfully for incoming socket."; + + OnIncomingConnection(upgrade_service_id); +} + +void ReconnectManager::BluetoothImpl::StopListeningForIncomingConnections() { + bluetooth_medium_.StopAcceptingConnections(reconnect_service_id_); +} + +bool ReconnectManager::BluetoothImpl::ConnectOverMedium() { + std::optional remote_mac_address = + client_->GetBluetoothMacAddress(endpoint_id_); + if (!remote_mac_address.has_value()) { + NEARBY_LOGS(INFO) + << "ReconnectBluetooth failed since remoteMacAddress is empty"; + return false; + } + auto& bluetooth_medium = mediums_->GetBluetoothClassic(); + BluetoothDevice remote_bluetooth_device = + bluetooth_medium.GetRemoteDevice(remote_mac_address.value()); + if (!remote_bluetooth_device.IsValid()) { + NEARBY_LOGS(INFO) + << "ReconnectBluetooth failed since remoteBluetoothDevice is null: " + << remote_mac_address.value(); + return false; + } + + bluetooth_socket_ = + bluetooth_medium.Connect(remote_bluetooth_device, reconnect_service_id_, + client_->GetCancellationFlag(endpoint_id_)); + + if (!bluetooth_socket_.IsValid()) { + NEARBY_LOGS(ERROR) << "Failed to reconnect to Bluetooth device " + << remote_bluetooth_device.GetName() + << " for endpoint(id=" << endpoint_id_ << ")."; + return false; + } + + reconnect_channel_ = std::make_unique( + UnWrapInitiatorReconnectServiceId(reconnect_service_id_), + /*channel_name=*/endpoint_id_, bluetooth_socket_); + if (reconnect_channel_ == nullptr) { + NEARBY_LOGS(ERROR) << "ReconnectBluetooth Failed to get the Bluetooth " + "channel, please retry "; + bluetooth_socket_.Close(); + return false; + } + return true; +} + +bool ReconnectManager::BluetoothImpl::SupportEncryptionDisabled() { + return false; +} + +void ReconnectManager::BluetoothImpl::QuietlyCloseChannelAndSocket() { + reconnect_channel_->Close(DisconnectionReason::UNFINISHED); + bluetooth_socket_.Close(); +} + +} // namespace connections +} // namespace nearby diff --git a/connections/implementation/reconnect_manager.h b/connections/implementation/reconnect_manager.h new file mode 100644 index 00000000..ecc29646 --- /dev/null +++ b/connections/implementation/reconnect_manager.h @@ -0,0 +1,237 @@ +// Copyright 2023 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef CORE_INTERNAL_RECONNECTION_MANAGER_H_ +#define CORE_INTERNAL_RECONNECTION_MANAGER_H_ + +#include +#include +#include + +#include "securegcm/ukey2_handshake.h" +#include "absl/base/thread_annotations.h" +#include "absl/container/flat_hash_map.h" +#include "absl/container/flat_hash_set.h" +#include "absl/functional/any_invocable.h" +#include "connections/implementation/client_proxy.h" +#include "connections/implementation/encryption_runner.h" +#include "connections/implementation/endpoint_channel.h" +#include "connections/implementation/endpoint_channel_manager.h" +#include "connections/implementation/mediums/bluetooth_classic.h" +#include "connections/implementation/mediums/mediums.h" +#include "connections/implementation/proto/offline_wire_formats.pb.h" +#include "internal/platform/bluetooth_classic.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/cancelable_alarm.h" +#include "internal/platform/count_down_latch.h" +#include "internal/platform/mutex.h" +#include "internal/platform/scheduled_executor.h" +#include "internal/platform/single_thread_executor.h" + +namespace nearby { +namespace connections { +using AutoReconnectFrame = ::location::nearby::connections::AutoReconnectFrame; +using OfflineFrame = ::location::nearby::connections::OfflineFrame; +using Medium = ::location::nearby::proto::connections::Medium; +using DisconnectionReason = + ::location::nearby::proto::connections::DisconnectionReason; + +class ReconnectManager { + public: + ReconnectManager(Mediums& mediums, EndpointChannelManager& channel_manager); + ~ReconnectManager(); + + struct AutoReconnectCallback { + absl::AnyInvocable + on_reconnect_success_cb; + absl::AnyInvocable + on_reconnect_failure_cb; + }; + + struct ReconnectMetadata { + ReconnectMetadata(bool is_incoming, AutoReconnectCallback callback, + bool send_disconnection_notification, + DisconnectionReason disconnection_reason, + const std::string& reconnect_service_id) + : reconnect_service_id(reconnect_service_id), + is_incoming(is_incoming), + send_disconnection_notification(send_disconnection_notification), + disconnection_reason(disconnection_reason) { + reconnect_cb = std::move(callback); + } + ~ReconnectMetadata() noexcept = default; + ReconnectMetadata(ReconnectMetadata&&) = default; + ReconnectMetadata& operator=(ReconnectMetadata&&) = default; + + AutoReconnectCallback reconnect_cb; + std::string reconnect_service_id; + bool is_incoming; + bool send_disconnection_notification; + DisconnectionReason disconnection_reason = + DisconnectionReason::UNKNOWN_DISCONNECTION_REASON; + }; + + // The entry point for AutoReconect, this API will do the auto reconnect for + // specified "endpoint_id" which connection was lost before. + bool AutoReconnect(ClientProxy* client, const std::string& endpoint_id, + AutoReconnectCallback& callback, + bool send_disconnection_notification, + DisconnectionReason disconnection_reason); + + private: + class MediumConnectionProcessor { + public: + virtual ~MediumConnectionProcessor() = default; + + virtual bool IsMediumRadioOn() const = 0; + virtual bool IsListeningForIncomingConnections() const = 0; + virtual bool StartListeningForIncomingConnections() = 0; + virtual void StopListeningForIncomingConnections() = 0; + virtual bool ConnectOverMedium() = 0; + virtual bool SupportEncryptionDisabled() = 0; + virtual void QuietlyCloseChannelAndSocket() = 0; + }; + + class BaseMediumImpl : public MediumConnectionProcessor { + public: + BaseMediumImpl(ClientProxy* client, const std::string& endpoint_id, + const std::string& reconnect_service_id, bool is_incoming, + Medium medium, Mediums* mediums, + EndpointChannelManager* channel_manager, + ReconnectManager& reconnect_manager) + : client_(client), + endpoint_id_(endpoint_id), + reconnect_service_id_(reconnect_service_id), + is_incoming_(is_incoming), + medium_(medium), + mediums_(mediums), + channel_manager_(channel_manager), + reconnect_manager_(reconnect_manager) {} + ~BaseMediumImpl() override = default; + + bool Run(); + + protected: + ClientProxy* client_; + std::string endpoint_id_; + std::string reconnect_service_id_; + bool is_incoming_; + Medium medium_ = Medium::UNKNOWN_MEDIUM; + Mediums* mediums_; + EndpointChannelManager* channel_manager_; + std::unique_ptr reconnect_channel_; + ReconnectManager& reconnect_manager_; + void OnIncomingConnection(const std::string& reconnect_service_id); + + private: + bool RehostForIncomingConnections(); + bool ReconnectToRemoteDevice(); + + std::string ReadClientIntroductionFrame(EndpointChannel* endpoint_channel); + bool ReadClientIntroductionAckFrame(EndpointChannel* endpoint_channel); + bool ReplaceChannelForEndpoint( + ClientProxy* client, const std::string& endpoint_id, + std::unique_ptr new_channel, + bool support_encryption_disabled, + absl::AnyInvocable stop_listening_incoming_connection); + EncryptionRunner::ResultListener GetResultListener(); + void OnEncryptionSuccessRunnable( + const std::string& endpoint_id, + std::unique_ptr ukey2, + const std::string& auth_token, const ByteArray& raw_auth_token); + void OnEncryptionFailureRunnable(const std::string& endpoint_id, + EndpointChannel* endpoint_channel); + void ProcessSuccessfulReconnection( + const std::string& endpoint_id, + absl::AnyInvocable stop_listening_incoming_connection); + void ProcessFailedReconnection( + const std::string& endpoint_id, + absl::AnyInvocable stop_listening_incoming_connection); + void StopListeningIfAllConnected( + const std::string& reconnect_service_id, + absl::AnyInvocable stop_listening_incoming_connection, + bool force_stop); + bool HasPendingIncomingConnections(const std::string& reconnect_service_id); + void CancelClearHostTimeoutAlarm(const std::string& service_id); + void ClearReconnectData(const std::string& service_id, bool is_incoming); + + std::unique_ptr wait_encryption_to_finish_; + bool replace_channel_succeed_; + }; + + class BluetoothImpl : public BaseMediumImpl { + public: + BluetoothImpl(ClientProxy* client_proxy, const std::string& endpoint_id, + const std::string& reconnect_service_id, bool is_incoming, + Medium medium, Mediums* mediums, + EndpointChannelManager* channel_manager, + ReconnectManager& reconnect_manager) + : BaseMediumImpl(client_proxy, endpoint_id, reconnect_service_id, + is_incoming, medium, mediums, channel_manager, + reconnect_manager), + bluetooth_medium_(mediums_->GetBluetoothClassic()) {} + + bool IsMediumRadioOn() const override; + bool IsListeningForIncomingConnections() const override; + bool StartListeningForIncomingConnections() override; + void StopListeningForIncomingConnections() override; + bool ConnectOverMedium() override; + bool SupportEncryptionDisabled() override; + void QuietlyCloseChannelAndSocket() override; + + private: + void OnIncomingBluetoothConnection(ClientProxy* client, + const std::string& upgrade_service_id, + BluetoothSocket socket); + + BluetoothClassic& bluetooth_medium_; + BluetoothSocket bluetooth_socket_; + }; + + bool Start(bool is_incoming, ClientProxy* client_proxy, + const std::string& endpoint_id, + const std::string& reconnect_service_id, Medium medium); + bool RunOnce(bool is_incoming, ClientProxy* client, + const std::string& endpoint_id, + const std::string& reconnect_service_id, Medium medium); + + void ClearReconnectData(ClientProxy* client, + const std::string& reconnect_service_id, + bool is_incoming); + void Shutdown(); + + Mediums* mediums_; + EndpointChannelManager* channel_manager_; + EncryptionRunner encryption_runner_; + + SingleThreadExecutor reconnect_executor_; + ScheduledExecutor alarm_executor_; + SingleThreadExecutor incoming_connection_cb_executor_; + SingleThreadExecutor encryption_cb_executor_; + + mutable RecursiveMutex mutex_; + absl::flat_hash_map> + listen_timeout_alarm_by_service_id_ ABSL_GUARDED_BY(mutex_); + absl::flat_hash_map> + new_endpoint_channels_; + absl::flat_hash_map endpoint_id_metadata_map_; + absl::flat_hash_set resumed_endpoints_; +}; + +} // namespace connections +} // namespace nearby + +#endif // CORE_INTERNAL_RECONNECTION_MANAGER_H_ diff --git a/connections/implementation/reconnect_manager_test.cc b/connections/implementation/reconnect_manager_test.cc new file mode 100644 index 00000000..1313fef3 --- /dev/null +++ b/connections/implementation/reconnect_manager_test.cc @@ -0,0 +1,149 @@ +// Copyright 2023 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "connections/implementation/reconnect_manager.h" + +#include +#include + +#include "gtest/gtest.h" +#include "absl/strings/string_view.h" +#include "absl/time/time.h" +#include "connections/implementation/client_proxy.h" +#include "connections/implementation/endpoint_channel_manager.h" +#include "connections/implementation/mediums/mediums.h" +#include "connections/implementation/simulation_user.h" +#include "connections/medium_selector.h" +#include "internal/platform/count_down_latch.h" +#include "internal/platform/logging.h" +#include "internal/platform/medium_environment.h" + +namespace nearby { +namespace connections { +namespace { + +constexpr absl::string_view kServiceId = "service-id"; +constexpr absl::string_view kDeviceA = "device-a"; +constexpr absl::string_view kDeviceB = "device-b"; +constexpr absl::Duration kDefaultTimeout = absl::Milliseconds(1000); + +constexpr BooleanMediumSelector kTestCases[] = { + BooleanMediumSelector{ + .bluetooth = true + }, +}; + +class ReconnectSimulatorUser : public SimulationUser { + public: + explicit ReconnectSimulatorUser( + absl::string_view name, + BooleanMediumSelector allowed = BooleanMediumSelector()) + : SimulationUser(std::string(name), allowed, + SetSafeToDisconnect(true, true, false, 3)) {} + ~ReconnectSimulatorUser() override { + NEARBY_LOGS(INFO) << "ReconnectSimulatorUser: [down] name=" << info_.data(); + } + + bool IsConnected() const { + return client_.IsConnectedToEndpoint(discovered_.endpoint_id); + } + + protected: +}; + +class ReconnectManagerTest + : public ::testing::TestWithParam { + protected: + bool SetupConnection(ReconnectSimulatorUser& user_a, + ReconnectSimulatorUser& user_b) { + user_a.StartAdvertising(std::string(kServiceId), &connection_latch_); + user_b.StartDiscovery(std::string(kServiceId), &discovery_latch_); + EXPECT_TRUE(discovery_latch_.Await(kDefaultTimeout).result()); + EXPECT_EQ(user_b.GetDiscovered().service_id, kServiceId); + EXPECT_EQ(user_b.GetDiscovered().endpoint_info, user_a.GetInfo()); + EXPECT_FALSE(user_b.GetDiscovered().endpoint_id.empty()); + NEARBY_LOGS(INFO) << "EP-B: [discovered]" + << user_b.GetDiscovered().endpoint_id; + user_b.RequestConnection(&connection_latch_); + EXPECT_TRUE(connection_latch_.Await(kDefaultTimeout).result()); + EXPECT_FALSE(user_a.GetDiscovered().endpoint_id.empty()); + NEARBY_LOGS(INFO) << "EP-A: [discovered]" + << user_a.GetDiscovered().endpoint_id; + NEARBY_LOGS(INFO) << "Both users discovered their peers."; + user_a.AcceptConnection(&accept_latch_); + user_b.AcceptConnection(&accept_latch_); + EXPECT_TRUE(accept_latch_.Await(kDefaultTimeout).result()); + NEARBY_LOG(INFO, "Both users reached connected state."); + return user_a.IsConnected() && user_b.IsConnected(); + } + + CountDownLatch discovery_latch_{1}; + CountDownLatch connection_latch_{2}; + CountDownLatch accept_latch_{2}; + CountDownLatch reject_latch_{1}; + MediumEnvironment& env_{MediumEnvironment::Instance()}; +}; + +TEST_P(ReconnectManagerTest, AllowReconnect) { + env_.Start(); + ReconnectSimulatorUser user_a(kDeviceA, GetParam()); + ReconnectSimulatorUser user_b(kDeviceB, GetParam()); + ASSERT_TRUE(SetupConnection(user_a, user_b)); + + Mediums mediums; + ReconnectManager::AutoReconnectCallback auto_reconnect_callback = { + .on_reconnect_success_cb = + [&](ClientProxy* client, const std::string& endpoint_id) { + NEARBY_LOGS(INFO) + << " Reconnect successfully for endpoint_id: " << endpoint_id; + }, + .on_reconnect_failure_cb = + [&](ClientProxy* client, const std::string& endpoint_id, + bool send_disconnection_notification, + DisconnectionReason disconnection_reason) { + NEARBY_LOGS(INFO) + << " Reconnect failed for endpoint_id: " << endpoint_id; + }, + }; + + auto& client_a = user_a.GetClient(); + auto& client_b = user_b.GetClient(); + EndpointChannelManager& ecm_a = user_a.GetEndpointChannelManager(); + EndpointChannelManager& ecm_b = user_b.GetEndpointChannelManager(); + + auto reconnect_manager_a = std::make_unique(mediums, ecm_a); + auto reconnect_manager_b = std::make_unique(mediums, ecm_b); + EXPECT_TRUE(reconnect_manager_a->AutoReconnect( + &client_a, user_a.GetDiscovered().endpoint_id, auto_reconnect_callback, + /*send_disconnection_notification=*/false, + DisconnectionReason::UNFINISHED)); + EXPECT_TRUE(reconnect_manager_b->AutoReconnect( + &client_b, user_b.GetDiscovered().endpoint_id, auto_reconnect_callback, + /*send_disconnection_notification=*/false, + DisconnectionReason::UNFINISHED)); + + NEARBY_LOGS(INFO) << "Test completed."; + user_a.Stop(); + user_b.Stop(); + env_.Stop(); +} + +INSTANTIATE_TEST_SUITE_P(ParametrisedReconnectManagerTest, ReconnectManagerTest, + ::testing::ValuesIn(kTestCases)); + +// More test will be added later. + +} // namespace +} // namespace connections +} // namespace nearby diff --git a/connections/implementation/service_id_constants.h b/connections/implementation/service_id_constants.h index 42728228..dace640c 100644 --- a/connections/implementation/service_id_constants.h +++ b/connections/implementation/service_id_constants.h @@ -19,6 +19,7 @@ #include "absl/strings/match.h" #include "absl/strings/string_view.h" +#include "absl/strings/strip.h" namespace nearby { namespace connections { @@ -28,6 +29,7 @@ constexpr absl::string_view kUnknownServiceId = "UNKNOWN_SERVICE"; // A suffix appended to service IDs when initiating a bandwidth upgrade to // distinguish the mediums from those used for advertising/discovery. constexpr absl::string_view kInitiatorUpgradeServiceIdPostfix = "_UPGRADE"; +constexpr absl::string_view kInitiatorReconnectServiceIdPostfix = "_RECONNECT"; // Returns true if |service_id| not empty and has the initiator's upgrade // postfix. @@ -47,6 +49,37 @@ inline std::string WrapInitiatorUpgradeServiceId(absl::string_view service_id) { std::string(kInitiatorUpgradeServiceIdPostfix); } +// Returns true if |service_id| not empty and has the initiator's reconnect +// postfix. +inline bool IsInitiatorReconnectServiceId(absl::string_view service_id) { + return !service_id.empty() && + absl::EndsWith(service_id, kInitiatorReconnectServiceIdPostfix); +} + +// Appends the kInitiatorReconnectServiceIdPostfix to |service_id| if necessary. +inline std::string WrapInitiatorReconnectServiceId( + absl::string_view service_id) { + // If |service_id| is empty or already has the reconnect postfix, do nothing. + if (service_id.empty() || IsInitiatorReconnectServiceId(service_id)) { + return std::string(service_id); + } + + return std::string(service_id) + + std::string(kInitiatorReconnectServiceIdPostfix); +} + +// Appends the kInitiatorReconnectServiceIdPostfix to |service_id| if necessary. +inline std::string UnWrapInitiatorReconnectServiceId( + absl::string_view service_id) { + // If |service_id| is empty or already has the reconnect postfix, do nothing. + if (service_id.empty() || !IsInitiatorReconnectServiceId(service_id)) { + return std::string(service_id); + } + + return std::string( + absl::StripSuffix(service_id, kInitiatorReconnectServiceIdPostfix)); +} + } // namespace connections } // namespace nearby diff --git a/connections/implementation/simulation_user.h b/connections/implementation/simulation_user.h index 78625612..9fba5c1b 100644 --- a/connections/implementation/simulation_user.h +++ b/connections/implementation/simulation_user.h @@ -15,6 +15,7 @@ #ifndef CORE_INTERNAL_SIMULATION_USER_H_ #define CORE_INTERNAL_SIMULATION_USER_H_ +#include #include #include @@ -45,13 +46,16 @@ namespace connections { class SetSafeToDisconnect { public: - explicit SetSafeToDisconnect(bool safe_to_disconnect, + explicit SetSafeToDisconnect(bool safe_to_disconnect, bool auto_reconnect, bool payload_received_ack, std::int32_t safe_to_disconnect_version) { NearbyFlags::GetInstance().OverrideBoolFlagValue( config_package_nearby::nearby_connections_feature:: kEnableSafeToDisconnect, safe_to_disconnect); + NearbyFlags::GetInstance().OverrideBoolFlagValue( + config_package_nearby::nearby_connections_feature::kEnableAutoReconnect, + auto_reconnect); NearbyFlags::GetInstance().OverrideBoolFlagValue( config_package_nearby::nearby_connections_feature:: kEnablePayloadReceivedAck, @@ -74,9 +78,10 @@ class SimulationUser { void Clear() { endpoint_id.clear(); } }; - explicit SimulationUser( - const std::string& device_name, - BooleanMediumSelector allowed = BooleanMediumSelector()) + SimulationUser(const std::string& device_name, + BooleanMediumSelector allowed = BooleanMediumSelector(), + SetSafeToDisconnect set_safe_to_disconnect = + SetSafeToDisconnect(true, false, true, 2)) : info_{ByteArray{device_name}}, advertising_options_{ { @@ -97,7 +102,8 @@ class SimulationUser { Strategy::kP2pCluster, allowed, }, - } {} + }, + set_safe_to_disconnect_(set_safe_to_disconnect) {} virtual ~SimulationUser() { Stop(); } void Stop() { pm_.DisconnectFromEndpointManager(); @@ -168,6 +174,9 @@ class SimulationUser { absl::AnyInvocable pred, absl::Duration timeout); + ClientProxy& GetClient() { return client_; } + EndpointChannelManager& GetEndpointChannelManager() { return ecm_; } + protected: // ConnectionListener callbacks void OnConnectionInitiated(const std::string& endpoint_id, @@ -206,7 +215,7 @@ class SimulationUser { AdvertisingOptions advertising_options_; ConnectionOptions connection_options_; DiscoveryOptions discovery_options_; - SetSafeToDisconnect set_safe_to_disconnect_{true, true, 2}; + SetSafeToDisconnect set_safe_to_disconnect_; ClientProxy client_; EndpointChannelManager ecm_; EndpointManager em_{&ecm_}; diff --git a/internal/platform/feature_flags.h b/internal/platform/feature_flags.h index 736c8015..43d602e1 100644 --- a/internal/platform/feature_flags.h +++ b/internal/platform/feature_flags.h @@ -63,6 +63,12 @@ class FeatureFlags { // on Windows when connecting to FP service id but the rfcomm is successful. bool skip_service_discovery_before_connecting_to_rfcomm = false; std::int32_t min_nc_version_supports_safe_to_disconnect = 1; + std::int32_t min_nc_version_supports_auto_reconnect = 3; + absl::Duration auto_reconnect_retry_delay_millis = absl::Milliseconds(5000); + absl::Duration auto_reconnect_timeout_millis = absl::Milliseconds(30000); + std::int32_t auto_reconnect_retry_attempts = 3; + absl::Duration auto_reconnect_skip_duplicated_endpoint_duration = + absl::Milliseconds(4000); // Android code won't be able to launch "payload_received_ack" feature for // in near future, so change "payload_received_ack" version from "2" to "5" // after auto-reconnect and auto-resume. @@ -73,6 +79,8 @@ class FeatureFlags { absl::Milliseconds(30000); absl::Duration safe_to_disconnect_remote_disc_delay_millis = absl::Milliseconds(10000); + absl::Duration safe_to_disconnect_auto_resume_timeout_millis = + absl::Milliseconds(60000); // If the receiver doesn't ack with payload_received_ack frame in 1s, the // sender will timeout the waiting. absl::Duration wait_payload_received_ack_millis = absl::Milliseconds(1000); diff --git a/internal/platform/implementation/BUILD b/internal/platform/implementation/BUILD index b672f4b0..3b11e279 100644 --- a/internal/platform/implementation/BUILD +++ b/internal/platform/implementation/BUILD @@ -40,6 +40,7 @@ cc_library( "timer.h", ], visibility = [ + "//connections/implementation:__subpackages__", "//connections/implementation/analytics:__subpackages__", "//fastpair:__subpackages__", "//internal/account:__pkg__",