diff --git a/connections/core.cc b/connections/core.cc index 827715cc..db156133 100644 --- a/connections/core.cc +++ b/connections/core.cc @@ -137,11 +137,12 @@ void Core::RequestConnection(absl::string_view endpoint_id, << "Client request connection with keep-alive frame as interval=" << connection_options.keep_alive_interval_millis << ", timeout=" << connection_options.keep_alive_timeout_millis - << ", which is un-expected. Change to default.", - connection_options.keep_alive_interval_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_interval_millis; + << ", which is un-expected. Change to default."; + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); + connection_options.keep_alive_interval_millis = + flags.keep_alive_interval_millis; connection_options.keep_alive_timeout_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_timeout_millis; + flags.keep_alive_timeout_millis; } router_->RequestConnection(&client_, endpoint_id, info, connection_options, @@ -404,10 +405,11 @@ void Core::RequestConnectionV3(const NearbyDevice& local_device, << connection_options.keep_alive_interval_millis << ", timeout=" << connection_options.keep_alive_timeout_millis << ", which is un-expected. Change to default."; + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); connection_options.keep_alive_interval_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_interval_millis; + flags.keep_alive_interval_millis; connection_options.keep_alive_timeout_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_timeout_millis; + flags.keep_alive_timeout_millis; } router_->RequestConnectionV3(&client_, remote_device, std::move(info), connection_options, std::move(result_cb)); @@ -437,10 +439,11 @@ void Core::RequestConnectionV3(const NearbyDevice& remote_device, << connection_options.keep_alive_interval_millis << ", timeout=" << connection_options.keep_alive_timeout_millis << ", which is un-expected. Change to default."; + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); connection_options.keep_alive_interval_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_interval_millis; + flags.keep_alive_interval_millis; connection_options.keep_alive_timeout_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_timeout_millis; + flags.keep_alive_timeout_millis; } router_->RequestConnectionV3(&client_, remote_device, std::move(info), connection_options, std::move(result_cb)); diff --git a/connections/implementation/base_endpoint_channel.h b/connections/implementation/base_endpoint_channel.h index 0924c794..9946b9fa 100644 --- a/connections/implementation/base_endpoint_channel.h +++ b/connections/implementation/base_endpoint_channel.h @@ -158,7 +158,7 @@ class BaseEndpointChannel : public EndpointChannel { // An encryptor/decryptor. May be null. mutable Mutex crypto_mutex_; std::shared_ptr crypto_context_ - ABSL_GUARDED_BY(crypto_mutex_) ABSL_PT_GUARDED_BY(crypto_mutex_); + ABSL_GUARDED_BY(crypto_mutex_); mutable Mutex is_paused_mutex_; ConditionVariable is_paused_cond_{&is_paused_mutex_}; diff --git a/connections/implementation/base_pcp_handler.cc b/connections/implementation/base_pcp_handler.cc index 3108f357..406d35f2 100644 --- a/connections/implementation/base_pcp_handler.cc +++ b/connections/implementation/base_pcp_handler.cc @@ -2052,11 +2052,12 @@ Exception BasePcpHandler::OnIncomingConnection( LOG(WARNING) << "Incoming connection has wrong keep-alive frame interval=" << connection_options.keep_alive_interval_millis << ", timeout=" << connection_options.keep_alive_timeout_millis - << " values; correct them as default.", - connection_options.keep_alive_interval_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_interval_millis; + << " values; correct them as default."; + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); + connection_options.keep_alive_interval_millis = + flags.keep_alive_interval_millis; connection_options.keep_alive_timeout_millis = - FeatureFlags::GetInstance().GetFlags().keep_alive_timeout_millis; + flags.keep_alive_timeout_millis; } const MediumMetadata& medium_metadata = connection_request.medium_metadata(); diff --git a/connections/implementation/bwu_manager.cc b/connections/implementation/bwu_manager.cc index 8f93a3de..fd41cfb2 100644 --- a/connections/implementation/bwu_manager.cc +++ b/connections/implementation/bwu_manager.cc @@ -83,22 +83,19 @@ BwuManager::BwuManager( mediums_(&mediums), endpoint_manager_(&endpoint_manager), channel_manager_(&channel_manager) { + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); if (config_.bandwidth_upgrade_retry_delay == absl::ZeroDuration()) { - if (FeatureFlags::GetInstance().GetFlags().use_exp_backoff_in_bwu_retry) { + if (flags.use_exp_backoff_in_bwu_retry) { config_.bandwidth_upgrade_retry_delay = - FeatureFlags::GetInstance() - .GetFlags() - .bwu_retry_exp_backoff_initial_delay; + flags.bwu_retry_exp_backoff_initial_delay; } else { config_.bandwidth_upgrade_retry_delay = absl::Seconds(5); } } if (config_.bandwidth_upgrade_retry_max_delay == absl::ZeroDuration()) { - if (FeatureFlags::GetInstance().GetFlags().use_exp_backoff_in_bwu_retry) { + if (flags.use_exp_backoff_in_bwu_retry) { config_.bandwidth_upgrade_retry_max_delay = - FeatureFlags::GetInstance() - .GetFlags() - .bwu_retry_exp_backoff_maximum_delay; + flags.bwu_retry_exp_backoff_maximum_delay; } else { config_.bandwidth_upgrade_retry_max_delay = absl::Seconds(10); } diff --git a/connections/implementation/bwu_manager_test.cc b/connections/implementation/bwu_manager_test.cc index cbae93aa..40cd3e50 100644 --- a/connections/implementation/bwu_manager_test.cc +++ b/connections/implementation/bwu_manager_test.cc @@ -120,6 +120,13 @@ class BwuManagerTest : public ::testing::Test { ~BwuManagerTest() override { bwu_manager_->Shutdown(); } + void SetSupportMultipleBwuMediums(bool support_multiple_bwu_mediums) { + FeatureFlags& feature_flags = FeatureFlags::GetMutableInstanceForTesting(); + FeatureFlags::Flags flags = feature_flags.GetFlags(); + flags.support_multiple_bwu_mediums = support_multiple_bwu_mediums; + feature_flags.SetFlags(flags); + } + // Create the initial device-to-device connection, before bandwidth upgrade. // Typically |medium| will be Bluetooth. FakeEndpointChannel* CreateInitialEndpoint(ClientProxy* client, @@ -315,8 +322,7 @@ class BwuManagerTestParam : public BwuManagerTest, public ::testing::WithParamInterface { protected: BwuManagerTestParam() { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = - GetParam(); + SetSupportMultipleBwuMediums(GetParam()); } }; @@ -475,7 +481,7 @@ TEST_P(BwuManagerTestParam, TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_MultipleEndpoints_FlagEnabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); // Say we have two already upgraded WebRTC connections for the same service. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -524,8 +530,7 @@ TEST_F(BwuManagerTest, TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_MultipleEndpoints_FlagDisabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = - false; + SetSupportMultipleBwuMediums(false); // Say we have two already upgraded WebRTC connections for the same service. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -581,7 +586,7 @@ TEST_F(BwuManagerTest, TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_MultipleServices_FlagEnabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); // Say we have two already upgraded WLAN connections for different services. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -635,8 +640,7 @@ TEST_F(BwuManagerTest, TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_MultipleServices_FlagDisabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = - false; + SetSupportMultipleBwuMediums(false); // Say we have two already upgraded WLAN connections for different services. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -694,7 +698,7 @@ TEST_F( BwuManagerTest, InitiateBwu_Revert_OnDisconnect_MultipleServicesAndEndpoints_FlagEnabled) { // Need support_multiple_bwu_mediums_ to run this test with multiple mediums. - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); // Say we have three upgraded connections for two different services and two // different mediums. @@ -843,7 +847,7 @@ TEST_F( } TEST_F(BwuManagerTest, InitiateBwu_Revert_OnUpgradeFailure_FlagEnabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); // Say we have two already upgraded WebRTC connections for service A. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -880,8 +884,7 @@ TEST_F(BwuManagerTest, InitiateBwu_Revert_OnUpgradeFailure_FlagEnabled) { } TEST_F(BwuManagerTest, InitiateBwu_Revert_OnUpgradeFailure_FlagDisabled) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = - false; + SetSupportMultipleBwuMediums(false); // Say we have two already upgraded WebRTC connections for service A. CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -917,7 +920,7 @@ TEST_F(BwuManagerTest, InitiateBwu_Revert_OnUpgradeFailure_FlagDisabled) { } TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_WifiDirect) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); OfflineFrame frame; CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -951,7 +954,7 @@ TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_WifiDirect) { } TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_Hotspot) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); @@ -981,7 +984,7 @@ TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_Hotspot) { } TEST_F(BwuManagerTest, InitiateBwu_Revert_OnDisconnect_Wlan) { - FeatureFlags::GetMutableFlagsForTesting().support_multiple_bwu_mediums = true; + SetSupportMultipleBwuMediums(true); CreateInitialEndpoint(&client_, kServiceIdA, kEndpointId1, Medium::BLUETOOTH); diff --git a/connections/implementation/endpoint_manager.cc b/connections/implementation/endpoint_manager.cc index 82677bca..ecde9b08 100644 --- a/connections/implementation/endpoint_manager.cc +++ b/connections/implementation/endpoint_manager.cc @@ -813,9 +813,8 @@ bool EndpointManager::ApplySafeToDisconnect(const std::string& endpoint_id, // TODO(b/303544913): clean up the safe-to-disconnect logic bool is_safe_disconnection = false; bool send_disconnection_frame = true; - absl::Duration timeout_millis = FeatureFlags::GetInstance() - .GetFlags() - .safe_to_disconnect_ack_delay_millis; + FeatureFlags::Flags flags = FeatureFlags::GetInstance().GetFlags(); + absl::Duration timeout_millis = flags.safe_to_disconnect_ack_delay_millis; bool is_wait_for_ack = true; switch (reason) { case DisconnectionReason::UPGRADED: @@ -832,9 +831,7 @@ bool EndpointManager::ApplySafeToDisconnect(const std::string& endpoint_id, case DisconnectionReason::REMOTE_DISCONNECTION: is_safe_disconnection = true; send_disconnection_frame = false; - timeout_millis = FeatureFlags::GetInstance() - .GetFlags() - .safe_to_disconnect_remote_disc_delay_millis; + timeout_millis = flags.safe_to_disconnect_remote_disc_delay_millis; is_wait_for_ack = false; break; default: diff --git a/internal/platform/BUILD b/internal/platform/BUILD index de9638d8..67f362f5 100644 --- a/internal/platform/BUILD +++ b/internal/platform/BUILD @@ -128,6 +128,7 @@ cc_library( ], deps = [ ":base", + "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_set", "@com_google_absl//absl/functional:any_invocable", "@com_google_absl//absl/synchronization", diff --git a/internal/platform/cancellation_flag.cc b/internal/platform/cancellation_flag.cc index 90b079bb..a83f546b 100644 --- a/internal/platform/cancellation_flag.cc +++ b/internal/platform/cancellation_flag.cc @@ -13,7 +13,10 @@ // limitations under the License. #include "internal/platform/cancellation_flag.h" +#include +#include "absl/container/flat_hash_set.h" +#include "absl/synchronization/mutex.h" #include "internal/platform/feature_flags.h" namespace nearby { @@ -28,7 +31,7 @@ CancellationFlag::CancellationFlag(bool cancelled) { } CancellationFlag::~CancellationFlag() { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); listeners_.clear(); } @@ -40,7 +43,7 @@ void CancellationFlag::Cancel() { absl::flat_hash_set listeners; { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); if (cancelled_) { // Someone already cancelled. Return immediately. return; @@ -62,14 +65,14 @@ void CancellationFlag::Uncancel() { } { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); assert(cancelled_); cancelled_ = false; } } bool CancellationFlag::Cancelled() const { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); // Return false as no-op if feature flag is not enabled. if (!FeatureFlags::GetInstance().GetFlags().enable_cancellation_flag) { @@ -80,13 +83,13 @@ bool CancellationFlag::Cancelled() const { } void CancellationFlag::RegisterOnCancelListener(CancelListener *listener) { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); listeners_.emplace(listener); } void CancellationFlag::UnregisterOnCancelListener(CancelListener *listener) { - absl::MutexLock lock(*mutex_.get()); + absl::MutexLock lock(*mutex_); listeners_.erase(listener); } diff --git a/internal/platform/cancellation_flag.h b/internal/platform/cancellation_flag.h index 9a8b7280..1e05539e 100644 --- a/internal/platform/cancellation_flag.h +++ b/internal/platform/cancellation_flag.h @@ -17,6 +17,7 @@ #include +#include "absl/base/thread_annotations.h" #include "absl/container/flat_hash_set.h" #include "absl/functional/any_invocable.h" #include "absl/synchronization/mutex.h" @@ -74,7 +75,7 @@ class CancellationFlag { ABSL_LOCKS_EXCLUDED(mutex_); int CancelListenersSize() const ABSL_LOCKS_EXCLUDED(mutex_) { - absl::MutexLock lock(mutex_.get()); + absl::MutexLock lock(*mutex_); return listeners_.size(); } diff --git a/internal/platform/feature_flags.h b/internal/platform/feature_flags.h index 02a574bc..cad1ba6a 100644 --- a/internal/platform/feature_flags.h +++ b/internal/platform/feature_flags.h @@ -126,13 +126,13 @@ class FeatureFlags { return *instance; } - const Flags& GetFlags() const ABSL_LOCKS_EXCLUDED(mutex_) { - absl::ReaderMutexLock lock(&mutex_); - return flags_; + static FeatureFlags& GetMutableInstanceForTesting() { + return const_cast(GetInstance()); } - static Flags& GetMutableFlagsForTesting() { - return const_cast(GetInstance()).flags_; + Flags GetFlags() const ABSL_LOCKS_EXCLUDED(mutex_) { + absl::ReaderMutexLock lock(mutex_); + return flags_; } // SetFlags for feature controlling diff --git a/internal/platform/feature_flags_test.cc b/internal/platform/feature_flags_test.cc index c4f24ee6..2ca92d43 100644 --- a/internal/platform/feature_flags_test.cc +++ b/internal/platform/feature_flags_test.cc @@ -28,8 +28,8 @@ constexpr FeatureFlags::Flags kTestFeatureFlags{ TEST(FeatureFlagsTest, CastUpdateWorks) { const FeatureFlags& features = FeatureFlags::GetInstance(); EXPECT_TRUE(features.GetFlags().enable_async_bandwidth_upgrade); - const_cast(FeatureFlags::GetInstance()) - .SetFlags({.enable_async_bandwidth_upgrade = false}); + FeatureFlags::GetMutableInstanceForTesting().SetFlags( + {.enable_async_bandwidth_upgrade = false}); EXPECT_FALSE(features.GetFlags().enable_async_bandwidth_upgrade); } diff --git a/internal/platform/medium_environment.cc b/internal/platform/medium_environment.cc index cbbf246a..f932bfc4 100644 --- a/internal/platform/medium_environment.cc +++ b/internal/platform/medium_environment.cc @@ -1147,7 +1147,7 @@ void MediumEnvironment::UnregisterWifiHotspotMedium( } void MediumEnvironment::SetFeatureFlags(const FeatureFlags::Flags& flags) { - const_cast(FeatureFlags::GetInstance()).SetFlags(flags); + FeatureFlags::GetMutableInstanceForTesting().SetFlags(flags); } std::optional MediumEnvironment::GetSimulatedClock() {