diff --git a/presence/data_types.h b/presence/data_types.h index 8b3fdb80..572926bc 100644 --- a/presence/data_types.h +++ b/presence/data_types.h @@ -19,6 +19,7 @@ #include #include "absl/functional/any_invocable.h" +#include "internal/platform/logging.h" #include "presence/presence_device.h" #include "presence/status.h" @@ -32,18 +33,29 @@ class ScanSession { ScanSession() : stop_scan_callback_( []() { return Status{Status::Value::kNotImplemented}; }) {} - explicit ScanSession(absl::AnyInvocable stop_scan_callback) + explicit ScanSession(absl::AnyInvocable stop_scan_callback) : stop_scan_callback_(std::move(stop_scan_callback)) {} + ~ScanSession() { StopScan(); } + Status StopScan() { - return stop_scan_callback_(); + if (stop_called_) { + NEARBY_LOGS(WARNING) << "StopScan already called."; + return Status{Status::Value::kError}; + } + stop_called_ = true; + if (stop_scan_callback_) { + return std::move(stop_scan_callback_)(); + } + return Status{Status::Value::kError}; } private: // Nearby library would provide the implementation of this callback in // runtime. Assigning with a default value NotImplemented to surface potential // issue where library failed to provide the implementation. - absl::AnyInvocable stop_scan_callback_; + absl::AnyInvocable stop_scan_callback_; + bool stop_called_ = false; }; // Callers would provide the implementation of these callbacks. If callers diff --git a/presence/implementation/scan_manager.cc b/presence/implementation/scan_manager.cc index ac9efdf5..ea7a7511 100644 --- a/presence/implementation/scan_manager.cc +++ b/presence/implementation/scan_manager.cc @@ -47,7 +47,8 @@ using ScanningCallback = namespace nearby { namespace presence { -ScanSession ScanManager::StartScan(ScanRequest scan_request, ScanCallback cb) { +std::unique_ptr ScanManager::StartScan(ScanRequest scan_request, + ScanCallback cb) { absl::BitGen gen; uint64_t id = absl::uniform_int_distribution(0, UINT64_MAX)(gen); ScanningCallback callback = ScanningCallback{ @@ -67,6 +68,11 @@ ScanSession ScanManager::StartScan(ScanRequest scan_request, ScanCallback cb) { [this](BlePeripheral& peripheral, BleAdvertisementData data) { NotifyFoundBle(data, peripheral); }}; + std::unique_ptr scanning_session = + mediums_->GetBle().StartScanning(scan_request, std::move(callback)); + if (scanning_session == nullptr) { + return nullptr; + } // We will not be needing the start_scan_cb anymore, so cb is ok to use here. AddScanCallback(id, MapElement{ .request = scan_request, @@ -74,9 +80,7 @@ ScanSession ScanManager::StartScan(ScanRequest scan_request, ScanCallback cb) { .decoder = AdvertisementDecoder(credential_manager_, scan_request), }); - std::unique_ptr scanning_session = - mediums_->GetBle().StartScanning(scan_request, std::move(callback)); - auto modified_scanning_session = ScanSession( + return std::make_unique( /*stop_scan_callback=*/[scanning_session_internal = std::move(scanning_session), this, id]() { @@ -94,7 +98,6 @@ ScanSession ScanManager::StartScan(ScanRequest scan_request, ScanCallback cb) { } return Status{.value = Status::Value::kSuccess}; }); - return modified_scanning_session; } void ScanManager::NotifyFoundBle(BleAdvertisementData data, diff --git a/presence/implementation/scan_manager.h b/presence/implementation/scan_manager.h index d33fd7f9..6de0953f 100644 --- a/presence/implementation/scan_manager.h +++ b/presence/implementation/scan_manager.h @@ -42,7 +42,8 @@ class ScanManager { } ~ScanManager() = default; - ScanSession StartScan(ScanRequest scan_request, ScanCallback cb) + std::unique_ptr StartScan(ScanRequest scan_request, + ScanCallback cb) ABSL_LOCKS_EXCLUDED(mutex_); // Below functions are test only. // Reference: go/totw/135#augmenting-the-public-api-for-tests diff --git a/presence/implementation/scan_manager_test.cc b/presence/implementation/scan_manager_test.cc index 64b9941a..72506ae8 100644 --- a/presence/implementation/scan_manager_test.cc +++ b/presence/implementation/scan_manager_test.cc @@ -16,6 +16,7 @@ #include +#include #include #include #include @@ -126,12 +127,12 @@ TEST_F(ScanManagerTest, CanStartThenStopScanning) { StartAdvertisingOn(ble2); // Start scanning - ScanSession scan_session = + auto scan_session = manager.StartScan(MakeDefaultScanRequest(), MakeDefaultScanCallback()); EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 1); EXPECT_TRUE(start_latch_.Await().Ok()); EXPECT_TRUE(found_latch_.Await().Ok()); - EXPECT_TRUE(scan_session.StopScan().Ok()); + EXPECT_TRUE(scan_session->StopScan().Ok()); EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 0); } @@ -147,9 +148,9 @@ TEST_F(ScanManagerTest, CannotStopScanTwice) { // Ensure that we have started scanning before we try to stop. env_.Sync(); NEARBY_LOGS(INFO) << "Stop scan"; - EXPECT_TRUE(scan_session.StopScan().Ok()); + EXPECT_TRUE(scan_session->StopScan().Ok()); NEARBY_LOGS(INFO) << "Stop scan again"; - EXPECT_FALSE(scan_session.StopScan().Ok()); + EXPECT_FALSE(scan_session->StopScan().Ok()); } TEST_F(ScanManagerTest, TestNoFilter) { @@ -164,14 +165,14 @@ TEST_F(ScanManagerTest, TestNoFilter) { // Start scanning ScanRequest scan_request_no_filter = MakeDefaultScanRequest(); scan_request_no_filter.scan_filters.clear(); - ScanSession scan_session = + auto scan_session = manager.StartScan(scan_request_no_filter, MakeDefaultScanCallback()); ASSERT_EQ(manager.ScanningCallbacksLengthForTest(), 1); ASSERT_TRUE(mediums.GetBle().IsAvailable()); EXPECT_TRUE(start_latch_.Await().Ok()); EXPECT_TRUE(found_latch_.Await().Ok()); - EXPECT_TRUE(scan_session.StopScan().Ok()); + EXPECT_TRUE(scan_session->StopScan().Ok()); EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 0); } @@ -200,7 +201,7 @@ TEST_F(ScanManagerTest, StopOneSessionFromAnotherDeadlock) { }; // we use scan_request_mismatch so this session's discovery doesn't get // triggered. - ScanSession scan_session = + auto scan_session = manager.StartScan(scan_request_mismatch, MakeDefaultScanCallback()); ScanCallback scanning_callback2 = { .start_scan_cb = @@ -211,11 +212,12 @@ TEST_F(ScanManagerTest, StopOneSessionFromAnotherDeadlock) { }, .on_discovered_cb = [&found_latch2, &scan_session](PresenceDevice pd) { + NEARBY_LOGS(INFO) << "scansession2 found"; found_latch2.CountDown(); - scan_session.StopScan(); + scan_session->StopScan(); }}; - ScanSession scan_session2 = manager.StartScan(MakeDefaultScanRequest(), - std::move(scanning_callback2)); + auto scan_session2 = manager.StartScan(MakeDefaultScanRequest(), + std::move(scanning_callback2)); ASSERT_EQ(manager.ScanningCallbacksLengthForTest(), 2); @@ -228,11 +230,53 @@ TEST_F(ScanManagerTest, StopOneSessionFromAnotherDeadlock) { EXPECT_TRUE(found_latch2.Await(absl::Milliseconds(1500)).result()); EXPECT_FALSE(found_latch_.Await(absl::Milliseconds(1500)).result()); // Session was stopped before, this should not be able to stop successfully. - EXPECT_FALSE(scan_session.StopScan().Ok()); + EXPECT_FALSE(scan_session->StopScan().Ok()); EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 1); - EXPECT_TRUE(scan_session2.StopScan().Ok()); + EXPECT_TRUE(scan_session2->StopScan().Ok()); EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 0); } + +TEST_F(ScanManagerTest, StopWhenScopeEnds) { + Mediums mediums; + ScanManager manager(mediums, credential_manager_); + ScanCallback scanning_callback = ScanCallback{ + .start_scan_cb = + [this](Status status) { + if (status.Ok()) { + start_latch_.CountDown(); + } + }, + }; + { + auto scan_session = manager.StartScan(MakeDefaultScanRequest(), + std::move(scanning_callback)); + EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 1); + // Ensure that we start scanning before we go out of scope. + env_.Sync(); + } + EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 0); +} + +TEST_F(ScanManagerTest, MoveDoesNotTriggerDestructor) { + Mediums mediums; + ScanManager manager(mediums, credential_manager_); + ScanCallback scanning_callback = ScanCallback{ + .start_scan_cb = + [this](Status status) { + if (status.Ok()) { + start_latch_.CountDown(); + } + }, + }; + auto scan_session = + manager.StartScan(MakeDefaultScanRequest(), std::move(scanning_callback)); + EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 1); + env_.Sync(); + auto scan_session_moved = std::move(scan_session); + // Make sure we don't trigger the destructor. + EXPECT_EQ(manager.ScanningCallbacksLengthForTest(), 1); +} + } // namespace } // namespace presence } // namespace nearby diff --git a/presence/implementation/service_controller_impl.cc b/presence/implementation/service_controller_impl.cc index 65a88acf..0875df75 100644 --- a/presence/implementation/service_controller_impl.cc +++ b/presence/implementation/service_controller_impl.cc @@ -21,8 +21,7 @@ namespace presence { std::unique_ptr ServiceControllerImpl::StartScan( ScanRequest scan_request, ScanCallback callback) { - return std::make_unique( - scan_manager_.StartScan(scan_request, callback)); + return scan_manager_.StartScan(scan_request, callback); } std::unique_ptr ServiceControllerImpl::StartBroadcast( BroadcastRequest broadcast_request, BroadcastCallback callback) {