diff --git a/sharing/nearby_connections_manager_impl.cc b/sharing/nearby_connections_manager_impl.cc index 78d1b9db..8339895f 100644 --- a/sharing/nearby_connections_manager_impl.cc +++ b/sharing/nearby_connections_manager_impl.cc @@ -424,7 +424,9 @@ void NearbyConnectionsManagerImpl::Connect( // Setup transfer manager. if (IsTransportTypeFlagsSet(transport_type, TransportType::kHighQuality)) { transfer_managers_[endpoint_id] = std::make_unique( - connections_callback_task_runner_, endpoint_id); + connections_callback_task_runner_, endpoint_id, + absl::bind_front(&NearbyConnectionsManagerImpl::SendWithoutDelay, + this)); } } @@ -499,23 +501,18 @@ void NearbyConnectionsManagerImpl::Send( RegisterPayloadStatusListener(payload->id, listener); } - if (transfer_managers_.contains(endpoint_id) && payload->content.is_file()) { - VLOG(1) << __func__ << ": Send payload " << payload->id << " to " - << endpoint_id << " to transfer manager. payload is file: " - << payload->content.is_file() << ", is bytes " - << payload->content.is_bytes(); - transfer_managers_.at(endpoint_id) - ->Send([&, endpoint_id = std::string(endpoint_id), - payload_copy = *payload]() { - VLOG(1) << __func__ << ": Send payload " << payload_copy.id << " to " - << endpoint_id; - auto sent_payload = std::make_unique(payload_copy); - SendWithoutDelay(endpoint_id, std::move(sent_payload)); - }); - transfer_managers_.at(endpoint_id)->StartTransfer(); - return; + if (payload->content.is_file()) { + const auto& it = transfer_managers_.find(endpoint_id); + if (it != transfer_managers_.end()) { + VLOG(1) << __func__ << ": Send payload " << payload->id << " to " + << endpoint_id << " to transfer manager. payload is file: " + << payload->content.is_file() << ", is bytes " + << payload->content.is_bytes(); + it->second->Send(std::move(payload)); + it->second->StartTransfer(); + return; + } } - SendWithoutDelay(endpoint_id, std::move(payload)); } @@ -740,9 +737,10 @@ void NearbyConnectionsManagerImpl::OnDisconnected( absl::string_view endpoint_id) { MutexLock lock(&mutex_); // Remove transfer manager. - if (transfer_managers_.contains(endpoint_id)) { - transfer_managers_[endpoint_id]->CancelTransfer(); - transfer_managers_.erase(endpoint_id); + const auto& transfer_manager_it = transfer_managers_.find(endpoint_id); + if (transfer_manager_it != transfer_managers_.end()) { + transfer_manager_it->second->CancelTransfer(); + transfer_managers_.erase(transfer_manager_it); } Status connection_layer_status = Status::kUnknown; @@ -772,8 +770,9 @@ void NearbyConnectionsManagerImpl::OnBandwidthChanged( << ": Bandwidth changed to medium=" << static_cast(medium) << "; endpoint_id=" << endpoint_id; - if (transfer_managers_.contains(endpoint_id)) { - transfer_managers_[endpoint_id]->OnMediumQualityChanged(medium); + const auto& transfer_manager_it = transfer_managers_.find(endpoint_id); + if (transfer_manager_it != transfer_managers_.end()) { + transfer_manager_it->second->OnMediumQualityChanged(medium); } current_upgraded_mediums_.insert_or_assign(endpoint_id, medium); diff --git a/sharing/transfer_manager.cc b/sharing/transfer_manager.cc index 863c970e..fa38384b 100644 --- a/sharing/transfer_manager.cc +++ b/sharing/transfer_manager.cc @@ -14,12 +14,12 @@ #include "sharing/transfer_manager.h" -#include #include #include -#include +#include #include "absl/base/nullability.h" +#include "absl/functional/any_invocable.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" #include "internal/platform/task_runner.h" @@ -43,28 +43,32 @@ bool IsHighQualityMedium(Medium medium) { } // namespace -TransferManager::TransferManager(TaskRunner* absl_nonnull runner, - absl::string_view endpoint_id) - : runner_(*runner), endpoint_id_(endpoint_id) {} +TransferManager::TransferManager( + TaskRunner* absl_nonnull runner, absl::string_view endpoint_id, + absl::AnyInvocable payload)> + deferred_send_function) + : runner_(*runner), + endpoint_id_(endpoint_id), + deferred_send_function_(std::move(deferred_send_function)) {} TransferManager::~TransferManager() { absl::MutexLock lock(mutex_); timeout_timer_.reset(); - pending_tasks_.clear(); } -void TransferManager::Send(std::function task) { +void TransferManager::Send(std::unique_ptr payload) { absl::MutexLock lock(mutex_); if (is_waiting_for_high_quality_medium_) { LOG(INFO) << "Connection to endpoint " << endpoint_id_ << " is waiting for a high quality medium, delaying payload transfer."; - pending_tasks_.push_back(task); + pending_payloads_.push(std::move(payload)); return; } - task(); + deferred_send_function_(endpoint_id_, std::move(payload)); } void TransferManager::OnMediumQualityChanged(Medium current_medium) { @@ -134,12 +138,13 @@ void TransferManager::StopWaitingForHighQualityMedium() { timeout_timer_.reset(); is_waiting_for_high_quality_medium_ = false; - for (const auto& task : pending_tasks_) { - LOG(INFO) << "Sending delayed payload to endpoint " << endpoint_id_; - task(); + LOG(INFO) << "Sending " << pending_payloads_.size() + << " delayed payloads to endpoint " << endpoint_id_; + while (!pending_payloads_.empty()) { + auto payload = std::move(pending_payloads_.front()); + pending_payloads_.pop(); + deferred_send_function_(endpoint_id_, std::move(payload)); } - - pending_tasks_.clear(); } } // namespace sharing diff --git a/sharing/transfer_manager.h b/sharing/transfer_manager.h index e19f248c..619b2b0e 100644 --- a/sharing/transfer_manager.h +++ b/sharing/transfer_manager.h @@ -15,13 +15,13 @@ #ifndef THIRD_PARTY_NEARBY_SHARING_TRANSFER_MANAGER_H_ #define THIRD_PARTY_NEARBY_SHARING_TRANSFER_MANAGER_H_ -#include #include +#include #include -#include #include "absl/base/nullability.h" #include "absl/base/thread_annotations.h" +#include "absl/functional/any_invocable.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" #include "absl/time/time.h" @@ -41,11 +41,14 @@ class TransferManager { static constexpr absl::Duration kMediumUpgradeTimeout = absl::Seconds(10); TransferManager(TaskRunner* absl_nonnull runner, - absl::string_view endpoint_id); + absl::string_view endpoint_id, + absl::AnyInvocable payload)> + deferred_send_function); ~TransferManager(); - void Send(std::function task) ABSL_LOCKS_EXCLUDED(mutex_); + void Send(std::unique_ptr payload) ABSL_LOCKS_EXCLUDED(mutex_); void OnMediumQualityChanged(Medium current_medium) ABSL_LOCKS_EXCLUDED(mutex_); bool StartTransfer() ABSL_LOCKS_EXCLUDED(mutex_); @@ -55,10 +58,14 @@ class TransferManager { void StopWaitingForHighQualityMedium() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); TaskRunner& runner_; - std::string endpoint_id_; + const std::string endpoint_id_; + absl::AnyInvocable payload)> + deferred_send_function_; absl::Mutex mutex_; bool is_waiting_for_high_quality_medium_ ABSL_GUARDED_BY(mutex_) = true; - std::vector> pending_tasks_ ABSL_GUARDED_BY(mutex_); + std::queue> pending_payloads_ + ABSL_GUARDED_BY(mutex_); std::unique_ptr timeout_timer_ ABSL_GUARDED_BY(mutex_) = nullptr; }; diff --git a/sharing/transfer_manager_test.cc b/sharing/transfer_manager_test.cc index f46869f5..63f0616c 100644 --- a/sharing/transfer_manager_test.cc +++ b/sharing/transfer_manager_test.cc @@ -14,6 +14,7 @@ #include "sharing/transfer_manager.h" +#include #include #include "gtest/gtest.h" @@ -37,11 +38,15 @@ TEST(TransferManager, MediumUpgradeSuccess) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -59,11 +64,15 @@ TEST(TransferManager, SendAfterMediumUpgradeSuccess) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -72,7 +81,7 @@ TEST(TransferManager, SendAfterMediumUpgradeSuccess) { notification.WaitForNotificationWithTimeout(kNotificationTimeout)); ASSERT_TRUE(is_called); is_called = false; - transfer_manager.Send([&]() { is_called = true; }); + transfer_manager.Send(std::make_unique()); ASSERT_TRUE(is_called); } @@ -82,11 +91,15 @@ TEST(TransferManager, MediumUpgradeFailed) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -102,11 +115,15 @@ TEST(TransferManager, MediumUpgradeTimeout) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -123,11 +140,15 @@ TEST(TransferManager, CancelStartedTransfer) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -145,11 +166,15 @@ TEST(TransferManager, CancelTimedOutMediumUpgrade) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); ASSERT_FALSE(is_called); ASSERT_TRUE(transfer_manager.StartTransfer()); @@ -167,11 +192,15 @@ TEST(TransferManager, MediumUpgradeBeforeStartTransfer) { absl::Notification notification; bool is_called = false; - TransferManager transfer_manager{&executor, kEndpointId}; - transfer_manager.Send([&]() { - is_called = true; - notification.Notify(); - }); + TransferManager transfer_manager{ + &executor, kEndpointId, + [&](absl::string_view endpoint_id, std::unique_ptr payload) { + is_called = true; + if (!notification.HasBeenNotified()) { + notification.Notify(); + } + }}; + transfer_manager.Send(std::make_unique()); transfer_manager.OnMediumQualityChanged(Medium::kWifiLan);