diff --git a/sharing/incoming_share_session.cc b/sharing/incoming_share_session.cc index e2a90a6c..65d1a398 100644 --- a/sharing/incoming_share_session.cc +++ b/sharing/incoming_share_session.cc @@ -42,12 +42,14 @@ #include "sharing/share_target.h" #include "sharing/text_attachment.h" #include "sharing/transfer_metadata.h" +#include "sharing/transfer_metadata_builder.h" #include "sharing/wifi_credentials_attachment.h" namespace nearby::sharing { namespace { using ::location::nearby::proto::sharing::OSType; +using ::nearby::sharing::service::proto::ConnectionResponseFrame; using ::nearby::sharing::service::proto::IntroductionFrame; using ::nearby::sharing::service::proto::V1Frame; using ::nearby::sharing::service::proto::WifiCredentials; @@ -175,7 +177,7 @@ bool IncomingShareSession::ProcessKeyVerificationResult( return true; } -void IncomingShareSession::RegisterPayloadListener( +void IncomingShareSession::AcceptTransfer( Clock* clock, NearbyConnectionsManager& connections_manager, std::function update_callback) { const absl::flat_hash_map& payload_map = @@ -196,6 +198,22 @@ void IncomingShareSession::RegisterPayloadListener( NL_VLOG(1) << __func__ << ": Accepted incoming files from share target - " << share_target().id; } + WriteResponseFrame(ConnectionResponseFrame::ACCEPT); + NL_VLOG(1) << __func__ << ": Successfully wrote response frame"; + + UpdateTransferMetadata( + TransferMetadataBuilder() + .set_status(TransferMetadata::Status::kAwaitingRemoteAcceptance) + .set_token(token()) + .build()); + + if (TryUpgradeBandwidth(connections_manager)) { + // Upgrade bandwidth regardless of advertising visibility because either + // the system or the user has verified the sender's identity; the + // stable identifiers potentially exposed by performing a bandwidth + // upgrade are no longer a concern. + NL_LOG(INFO) << __func__ << ": Upgrade bandwidth when sending accept."; + } } bool IncomingShareSession::UpdateFilePayloadPaths( @@ -345,4 +363,16 @@ std::vector IncomingShareSession::GetPayloadFilePaths() return file_paths; } +bool IncomingShareSession::TryUpgradeBandwidth( + NearbyConnectionsManager& connections_manager) { + if (!bandwidth_upgrade_requested_ && + attachment_container().GetTotalAttachmentsSize() >= + kAttachmentsSizeThresholdOverHighQualityMedium) { + connections_manager.UpgradeBandwidth(endpoint_id()); + bandwidth_upgrade_requested_ = true; + return true; + } + return false; +} + } // namespace nearby::sharing diff --git a/sharing/incoming_share_session.h b/sharing/incoming_share_session.h index b07449a6..dee0da31 100644 --- a/sharing/incoming_share_session.h +++ b/sharing/incoming_share_session.h @@ -70,7 +70,8 @@ class IncomingShareSession : public ShareSession { bool UpdateFilePayloadPaths( const NearbyConnectionsManager& connections_manager); - void RegisterPayloadListener( + // Accept the transfer and begin listening for payload transfer updates. + void AcceptTransfer( Clock* clock, NearbyConnectionsManager& connections_manager, std::function update_callback); @@ -82,6 +83,10 @@ class IncomingShareSession : public ShareSession { // Returns the file paths of all file payloads. std::vector GetPayloadFilePaths() const; + // Upgrade bandwidth if it is needed. + // Returns true if bandwidth upgrade was requested. + bool TryUpgradeBandwidth(NearbyConnectionsManager& connections_manager); + protected: void InvokeTransferUpdateCallback(const TransferMetadata& metadata) override; bool OnNewConnection(NearbyConnection* connection) override; @@ -93,6 +98,8 @@ class IncomingShareSession : public ShareSession { std::function transfer_update_callback_; + + bool bandwidth_upgrade_requested_ = false; }; } // namespace nearby::sharing diff --git a/sharing/incoming_share_session_test.cc b/sharing/incoming_share_session_test.cc index f3102e44..a8330b79 100644 --- a/sharing/incoming_share_session_test.cc +++ b/sharing/incoming_share_session_test.cc @@ -50,16 +50,21 @@ namespace nearby::sharing { namespace { using ::location::nearby::proto::sharing::OSType; +using ::nearby::sharing::service::proto::ConnectionResponseFrame; using ::nearby::sharing::service::proto::FileMetadata; +using ::nearby::sharing::service::proto::Frame; using ::nearby::sharing::service::proto::IntroductionFrame; using ::nearby::sharing::service::proto::TextMetadata; using ::nearby::sharing::service::proto::V1Frame; using ::nearby::sharing::service::proto::WifiCredentials; using ::nearby::sharing::service::proto::WifiCredentialsMetadata; +using ::testing::_; using ::testing::Eq; +using ::testing::Invoke; using ::testing::IsEmpty; using ::testing::IsFalse; using ::testing::IsTrue; +using ::testing::MockFunction; using ::testing::UnorderedElementsAre; constexpr absl::string_view kEndpointId = "ABCD"; @@ -95,7 +100,7 @@ class IncomingShareSessionTest : public ::testing::Test { protected: IncomingShareSessionTest() : session_(task_runner_, std::string(kEndpointId), share_target_, - [](const IncomingShareSession&, const TransferMetadata&) {}) { + transfer_metadata_callback_.AsStdFunction()) { NL_CHECK( proto2::TextFormat::ParseFromString(R"pb( file_metadata { @@ -149,6 +154,8 @@ class IncomingShareSessionTest : public ::testing::Test { FakeClock clock_; FakeTaskRunner task_runner_{&clock_, 1}; ShareTarget share_target_; + MockFunction + transfer_metadata_callback_; IncomingShareSession session_; IntroductionFrame introduction_frame_; }; @@ -564,14 +571,24 @@ TEST_F(IncomingShareSessionTest, FinalizePayloadsMissingWifiPayloads) { IsFalse()); } -TEST_F(IncomingShareSessionTest, RegisterPayloadListenerSuccess) { +TEST_F(IncomingShareSessionTest, AcceptTransferSuccess) { + NearbySharingDecoderImpl nearby_sharing_decoder; + FakeNearbyConnection connection; + EXPECT_TRUE( + session_.OnConnected(nearby_sharing_decoder, absl::Now(), &connection)); EXPECT_THAT(session_.ProcessIntroduction(introduction_frame_), Eq(std::nullopt)); + EXPECT_CALL(transfer_metadata_callback_, Call(_, _)) + .WillOnce(Invoke([](const IncomingShareSession& session, + const TransferMetadata& metadata) { + EXPECT_EQ(metadata.status(), + TransferMetadata::Status::kAwaitingRemoteAcceptance); + })); + FakeNearbyConnectionsManager connections_manager; FakeClock clock; - - session_.RegisterPayloadListener(&clock, connections_manager, - [](int64_t, TransferMetadata) {}); + session_.AcceptTransfer(&clock, connections_manager, + [](int64_t, TransferMetadata) {}); for (auto it : session_.attachment_payload_map()) { EXPECT_THAT( @@ -579,6 +596,13 @@ TEST_F(IncomingShareSessionTest, RegisterPayloadListenerSuccess) { .lock(), Eq(session_.payload_tracker().lock())); } + std::vector frame_data = connection.GetWrittenData(); + Frame frame; + ASSERT_TRUE(frame.ParseFromArray(frame_data.data(), frame_data.size())); + ASSERT_EQ(frame.version(), Frame::V1); + ASSERT_EQ(frame.v1().type(), V1Frame::RESPONSE); + EXPECT_EQ(frame.v1().connection_response().status(), + ConnectionResponseFrame::ACCEPT); } TEST_F(IncomingShareSessionTest, ProcessKeyVerificationResultSuccess) { @@ -722,5 +746,48 @@ TEST_F(IncomingShareSessionTest, ProcessKeyVerificationResultUnknown) { EXPECT_THAT(introduction_received, IsFalse()); } +TEST_F(IncomingShareSessionTest, TryUpgradeBandwidthNotNeeded) { + NearbySharingDecoderImpl decoder; + FakeNearbyConnection connection; + FakeNearbyConnectionsManager connections_manager; + session_.OnConnected(decoder, absl::Now(), &connection); + + EXPECT_THAT(session_.TryUpgradeBandwidth(connections_manager), IsFalse()); +} + +TEST_F(IncomingShareSessionTest, TryUpgradeBandwidthNeeded) { + IntroductionFrame introduction_frame; + NL_CHECK( + proto2::TextFormat::ParseFromString(R"pb( + file_metadata { + id: 1234 + size: 1000000 + name: "file_name1" + mime_type: "application/pdf" + type: DOCUMENT + parent_folder: "parent_folder1" + payload_id: 9876 + } + file_metadata { + id: 1235 + size: 200 + name: "file_name2" + mime_type: "image/jpeg" + type: IMAGE + parent_folder: "parent_folder2" + payload_id: 9875 + } + )pb", + &introduction_frame)); + NearbySharingDecoderImpl decoder; + FakeNearbyConnection connection; + FakeNearbyConnectionsManager connections_manager; + session_.OnConnected(decoder, absl::Now(), &connection); + EXPECT_THAT(session_.ProcessIntroduction(introduction_frame), + Eq(std::nullopt)); + + EXPECT_THAT(session_.TryUpgradeBandwidth(connections_manager), IsTrue()); +} + } // namespace } // namespace nearby::sharing diff --git a/sharing/nearby_sharing_service_impl.cc b/sharing/nearby_sharing_service_impl.cc index 8d8c6b13..f2f49101 100644 --- a/sharing/nearby_sharing_service_impl.cc +++ b/sharing/nearby_sharing_service_impl.cc @@ -846,14 +846,24 @@ void NearbySharingServiceImpl::Accept( ? std::get<2>(*metadata).status() : TransferMetadata::Status::kUnknown)) { NL_LOG(WARNING) << __func__ << ": out of order API call."; - status_codes_callback(StatusCodes::kOutOfOrderApiCall); + std::move(status_codes_callback)(StatusCodes::kOutOfOrderApiCall); return; } if (is_incoming) { IncomingShareSession* incoming_session = GetIncomingShareSession(share_target_id); - ReceivePayloads(*incoming_session, std::move(status_codes_callback)); + mutual_acceptance_timeout_alarm_.reset(); + + // Log analytics event of starting to receive payloads. + analytics_recorder_->NewReceiveAttachmentsStart( + incoming_session->session_id(), + incoming_session->attachment_container()); + incoming_session->AcceptTransfer( + context_->GetClock(), *nearby_connections_manager_, + absl::bind_front( + &NearbySharingServiceImpl::OnPayloadTransferUpdate, this)); + std::move(status_codes_callback)(StatusCodes::kOk); return; } @@ -2406,41 +2416,6 @@ void NearbySharingServiceImpl::OnTransferStarted(bool is_incoming) { InvalidateSurfaceState(); } -void NearbySharingServiceImpl::ReceivePayloads( - IncomingShareSession& session, - std::function status_codes_callback) { - mutual_acceptance_timeout_alarm_.reset(); - - // Log analytics event of starting to receive payloads. - analytics_recorder_->NewReceiveAttachmentsStart( - session.session_id(), session.attachment_container()); - session.RegisterPayloadListener( - context_->GetClock(), *nearby_connections_manager_, - absl::bind_front(&NearbySharingServiceImpl::OnPayloadTransferUpdate, - this)); - session.WriteResponseFrame( - nearby::sharing::service::proto::ConnectionResponseFrame::ACCEPT); - NL_VLOG(1) << __func__ << ": Successfully wrote response frame"; - - session.UpdateTransferMetadata( - TransferMetadataBuilder() - .set_status(TransferMetadata::Status::kAwaitingRemoteAcceptance) - .set_token(session.token()) - .build()); - - if (session.attachment_container().GetTotalAttachmentsSize() >= - kAttachmentsSizeThresholdOverHighQualityMedium) { - // Upgrade bandwidth regardless of advertising visibility because either - // the system or the user has verified the sender's identity; the - // stable identifiers potentially exposed by performing a bandwidth - // upgrade are no longer a concern. - NL_LOG(INFO) << __func__ << ": Upgrade bandwidth when receiving accept."; - nearby_connections_manager_->UpgradeBandwidth(session.endpoint_id()); - } - - std::move(status_codes_callback)(StatusCodes::kOk); -} - NearbySharingService::StatusCodes NearbySharingServiceImpl::SendPayloads( ShareSession& session) { NL_VLOG(1) << __func__ << ": Preparing to send payloads to " @@ -3040,12 +3015,10 @@ void NearbySharingServiceImpl::OnReceivedIntroduction( sharing::config_package_nearby::nearby_sharing_feature:: kUpgradeBandwidthAfterAccept)) { if (frame->has_start_transfer() && frame->start_transfer()) { - if (session->attachment_container().GetTotalAttachmentsSize() >= - kAttachmentsSizeThresholdOverHighQualityMedium) { + if (session->TryUpgradeBandwidth(*nearby_connections_manager_)) { NL_LOG(INFO) << __func__ << ": Upgrade bandwidth when receiving an introduction frame."; - nearby_connections_manager_->UpgradeBandwidth(session->endpoint_id()); } } } @@ -3286,22 +3259,28 @@ void NearbySharingServiceImpl::HandleProgressUpdateFrame( progress_update_frame) { if (progress_update_frame.has_start_transfer() && progress_update_frame.start_transfer()) { - ShareSession* session = GetShareSession(share_target_id); + IncomingShareSession* session = GetIncomingShareSession(share_target_id); - if (session != nullptr && - session->attachment_container().GetTotalAttachmentsSize() >= - kAttachmentsSizeThresholdOverHighQualityMedium) { + if (session == nullptr || !session->IsConnected()) { + NL_LOG(ERROR) << "Received ProgressUpdate Frame on unknown session"; + return; + } + NL_LOG(INFO) << __func__ << ": Received progress for ShareTarget " + << share_target_id << " : " + << progress_update_frame.progress(); + // TODO(b/338468927): Check if this is actually needed. + // Bandwidth upgrade was already requested in Accept. + if (session->TryUpgradeBandwidth(*nearby_connections_manager_)) { NL_LOG(INFO) << __func__ << ": Upgrade bandwidth when receiving progress update frame " "for endpoint " << session->endpoint_id(); - nearby_connections_manager_->UpgradeBandwidth(session->endpoint_id()); } } if (progress_update_frame.has_progress()) { - NL_LOG(INFO) << __func__ << ": Current progress for ShareTarget " + NL_VLOG(1) << __func__ << ": Current progress for ShareTarget " << share_target_id << " is " << progress_update_frame.progress(); } diff --git a/sharing/nearby_sharing_service_impl.h b/sharing/nearby_sharing_service_impl.h index d5fe248a..b45a728f 100644 --- a/sharing/nearby_sharing_service_impl.h +++ b/sharing/nearby_sharing_service_impl.h @@ -312,9 +312,6 @@ class NearbySharingServiceImpl void OnTransferComplete(); void OnTransferStarted(bool is_incoming); - void ReceivePayloads( - IncomingShareSession& session, - std::function status_codes_callback); StatusCodes SendPayloads(ShareSession& session); void OnOutgoingConnection(absl::Time connect_start_time,