diff --git a/sharing/BUILD b/sharing/BUILD index d602ea63..b4ed7d17 100644 --- a/sharing/BUILD +++ b/sharing/BUILD @@ -597,3 +597,17 @@ cc_test( "@com_google_googletest//:gtest_main", ], ) + +cc_test( + name = "share_target_info_test", + srcs = ["share_target_info_test.cc"], + deps = [ + ":nearby_sharing_service", + ":transfer_metadata", + ":types", + "//internal/platform/implementation/g3", # fixdeps: keep + "@com_github_protobuf_matchers//protobuf-matchers", + "@com_google_absl//absl/strings:string_view", + "@com_google_googletest//:gtest_main", + ], +) diff --git a/sharing/nearby_sharing_service_impl.cc b/sharing/nearby_sharing_service_impl.cc index 90c3b0ae..1d5999c8 100644 --- a/sharing/nearby_sharing_service_impl.cc +++ b/sharing/nearby_sharing_service_impl.cc @@ -297,6 +297,10 @@ void NearbySharingServiceImpl::Shutdown( }); } +bool NearbySharingServiceImpl::IsShuttingDown() { + return (is_shutting_down_ == nullptr || *is_shutting_down_); +} + void NearbySharingServiceImpl::Cleanup() { SetInHighVisibility(false); @@ -781,19 +785,13 @@ void NearbySharingServiceImpl::Reject( return; } - NearbyConnection* connection = info->connection(); - RunOnNearbySharingServiceThreadDelayed( "incoming_rejection_delay", kIncomingRejectionDelay, [this, share_target_id]() { CloseConnection(share_target_id); }); + // kRejected status already sent below, no need to send on disconnect. + info->set_disconnect_status(TransferMetadata::Status::kUnknown); - connection->SetDisconnectionListener([this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id]() { - UnregisterShareTarget(share_target_id); - }); - }); - + NearbyConnection* connection = info->connection(); WriteResponseFrame( *connection, nearby::sharing::service::proto::ConnectionResponseFrame::REJECT); @@ -883,16 +881,8 @@ void NearbySharingServiceImpl::DoCancel( NL_LOG(INFO) << "Disconnect fully established endpoint id:" << info->endpoint_id(); if (is_initiator_of_cancellation) { - info->connection()->SetDisconnectionListener( - [this, share_target_id, info]() { - info->set_connection(nullptr); - RunOnNearbySharingServiceThread( - "api_unregister_share_target", [this, share_target_id]() { - NL_LOG(INFO) - << "Unregister share target in disconnection listener."; - UnregisterShareTarget(share_target_id); - }); - }); + // kCancelled status already sent above, no need to send on disconnect. + info->set_disconnect_status(TransferMetadata::Status::kUnknown); RunOnNearbySharingServiceThreadDelayed( "initiator_cancel_delay", kInitiatorCancelDelay, @@ -1031,18 +1021,15 @@ void NearbySharingServiceImpl::OnIncomingConnection( ShareTargetInfo& share_target_info = GetOrCreateShareTargetInfo(placeholder_share_target, endpoint_id); share_target_info.set_connection(connection); + share_target_info.set_disconnect_status( + TransferMetadata::Status::kAwaitingRemoteAcceptanceFailed); + connection->SetDisconnectionListener([this, placeholder_share_target_id]() { + OnConnectionDisconnected(placeholder_share_target_id); + }); // Set receiving session id. receiving_session_id_ = analytics_recorder_->GenerateNextId(); - connection->SetDisconnectionListener([this, placeholder_share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, placeholder_share_target_id]() { - OnConnectionDisconnected( - placeholder_share_target_id, - TransferMetadata::Status::kAwaitingRemoteAcceptanceFailed); - }); - }); std::unique_ptr advertisement = decoder_->DecodeAdvertisement(endpoint_info); @@ -2382,6 +2369,7 @@ void NearbySharingServiceImpl::OnRotateBackgroundAdvertisementTimerFired() { void NearbySharingServiceImpl::RemoveOutgoingShareTargetWithEndpointId( absl::string_view endpoint_id) { + disconnection_timeout_alarms_.erase(endpoint_id); auto it = outgoing_share_target_map_.find(endpoint_id); if (it == outgoing_share_target_map_.end()) { return; @@ -2601,6 +2589,10 @@ void NearbySharingServiceImpl::OnOutgoingConnection( } info->set_connection(connection); + info->set_disconnect_status( + TransferMetadata::Status::kUnexpectedDisconnection); + connection->SetDisconnectionListener( + [this, share_target_id]() { OnConnectionDisconnected(share_target_id); }); // Log analytics event of establishing connection. analytics_recorder_->NewEstablishConnection( @@ -2614,14 +2606,6 @@ void NearbySharingServiceImpl::OnOutgoingConnection( : 0, std::nullopt); - connection->SetDisconnectionListener([this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id]() { - OnConnectionDisconnected( - share_target_id, - TransferMetadata::Status::kUnexpectedDisconnection); - }); - }); std::optional four_digit_token = TokenToFourDigitString( nearby_connections_manager_->GetRawAuthenticationToken( @@ -2999,12 +2983,7 @@ void NearbySharingServiceImpl::Fail(int64_t share_target_id, "incoming_rejection_delay", kIncomingRejectionDelay, [this, share_target_id]() { CloseConnection(share_target_id); }); - connection->SetDisconnectionListener([this, share_target_id, status]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id, status]() { - OnConnectionDisconnected(share_target_id, status); - }); - }); + info->set_disconnect_status(status); // Send response to remote device. nearby::sharing::service::proto::ConnectionResponseFrame::Status @@ -3258,15 +3237,11 @@ void NearbySharingServiceImpl::OnIncomingDecryptedCertificate( ShareTargetInfo* share_target_info = GetShareTargetInfo(share_target_id); NL_DCHECK(share_target_info); share_target_info->set_connection(connection); - - connection->SetDisconnectionListener([this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id]() { - OnConnectionDisconnected( - share_target_id, - TransferMetadata::Status::kAwaitingRemoteAcceptanceFailed); - }); - }); + share_target_info->set_disconnect_status( + TransferMetadata::Status::kAwaitingRemoteAcceptanceFailed); + // Need to rebind the disconnect listener to the new share target id. + connection->SetDisconnectionListener( + [this, share_target_id]() { OnConnectionDisconnected(share_target_id); }); std::optional four_digit_token = TokenToFourDigitString( nearby_connections_manager_->GetRawAuthenticationToken(endpoint_id)); @@ -3746,7 +3721,6 @@ void NearbySharingServiceImpl::OnStorageCheckCompleted( << share_target_id; return; } - NearbyConnection* connection = info->connection(); mutual_acceptance_timeout_alarm_->Stop(); mutual_acceptance_timeout_alarm_->Start( @@ -3781,14 +3755,8 @@ void NearbySharingServiceImpl::OnStorageCheckCompleted( return; } - connection->SetDisconnectionListener([this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id]() { - OnConnectionDisconnected( - share_target_id, - TransferMetadata::Status::kUnexpectedDisconnection); - }); - }); + info->set_disconnect_status( + TransferMetadata::Status::kUnexpectedDisconnection); auto* frames_reader = info->frames_reader(); if (!frames_reader) { @@ -3888,11 +3856,13 @@ void NearbySharingServiceImpl::HandleProgressUpdateFrame( } void NearbySharingServiceImpl::OnConnectionDisconnected( - int64_t share_target_id, TransferMetadata::Status status) { + int64_t share_target_id) { + if (IsShuttingDown()) { + return; + } ShareTargetInfo* info = GetShareTargetInfo(share_target_id); - if (info) { - info->UpdateTransferMetadata( - TransferMetadataBuilder().set_status(status).build()); + if (info != nullptr) { + info->OnDisconnect(); } UnregisterShareTarget(share_target_id); } @@ -4051,6 +4021,11 @@ void NearbySharingServiceImpl::OnPayloadTransferUpdate( .build() : metadata); + if (payload_incomplete || + TransferMetadata::IsFinalStatus(metadata.status())) { + // final status already sent, no need to send again on disconnect. + info->set_disconnect_status(TransferMetadata::Status::kUnknown); + } // Cancellation has its own disconnection strategy, possibly adding a // delay before disconnection to provide the other party time to process // the cancellation. @@ -4072,13 +4047,6 @@ bool NearbySharingServiceImpl::OnIncomingPayloadsComplete( return false; } - NearbyConnection* connection = info->connection(); - - connection->SetDisconnectionListener([this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", - [this, share_target_id]() { UnregisterShareTarget(share_target_id); }); - }); if (!update_file_paths_in_progress_) { UpdateFilePath(share_target); @@ -4267,18 +4235,7 @@ void NearbySharingServiceImpl::Disconnect(int64_t share_target_id, disconnection_timeout_alarms_[endpoint_id] = std::move(timer); - // Stop the disconnection timeout if the connection has been closed already. - if (share_target_info->connection()) { - share_target_info->connection()->SetDisconnectionListener( - [this, share_target_id, share_target_info, endpoint_id]() { - share_target_info->set_connection(nullptr); - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id, endpoint_id]() { - OnDisconnectingConnectionDisconnected(share_target_id, - endpoint_id); - }); - }); - } + share_target_info->set_disconnect_status(TransferMetadata::Status::kUnknown); } void NearbySharingServiceImpl::OnDisconnectingConnectionTimeout( @@ -4291,12 +4248,6 @@ void NearbySharingServiceImpl::OnDisconnectingConnectionTimeout( nearby_connections_manager_->Disconnect(endpoint_id); } -void NearbySharingServiceImpl::OnDisconnectingConnectionDisconnected( - int64_t share_target_id, absl::string_view endpoint_id) { - disconnection_timeout_alarms_.erase(endpoint_id); - UnregisterShareTarget(share_target_id); -} - ShareTargetInfo& NearbySharingServiceImpl::GetOrCreateShareTargetInfo( const ShareTarget& share_target, absl::string_view endpoint_id) { if (share_target.is_incoming) { @@ -4540,17 +4491,8 @@ void NearbySharingServiceImpl::AbortAndCloseConnectionIfNecessary( // Close connection if necessary. if (info->connection()) { - // Ensure that the disconnect listener is set to UnregisterShareTarget - // because the other listeners also try to record a final status - // metric. - info->connection()->SetDisconnectionListener( - [this, share_target_id]() { - RunOnNearbySharingServiceThread( - "disconnection_listener", [this, share_target_id]() { - UnregisterShareTarget(share_target_id); - }); - }); - + // Final status already sent above. No need to send it again. + info->set_disconnect_status(TransferMetadata::Status::kUnknown); info->connection()->Close(); } }); @@ -4658,7 +4600,7 @@ bool NearbySharingServiceImpl::ReadyToAccept( void NearbySharingServiceImpl::RunOnNearbySharingServiceThread( absl::string_view task_name, absl::AnyInvocable task) { - if (is_shutting_down_ == nullptr || *is_shutting_down_) { + if (IsShuttingDown()) { NL_LOG(WARNING) << __func__ << ": Skip the task " << task_name << " due to service is shutting down."; return; @@ -4689,7 +4631,7 @@ void NearbySharingServiceImpl::RunOnNearbySharingServiceThread( void NearbySharingServiceImpl::RunOnNearbySharingServiceThreadDelayed( absl::string_view task_name, absl::Duration delay, absl::AnyInvocable task) { - if (is_shutting_down_ == nullptr || *is_shutting_down_) { + if (IsShuttingDown()) { NL_LOG(WARNING) << __func__ << ": Skip the delayed task " << task_name << " due to service is shutting down."; return; @@ -4719,7 +4661,7 @@ void NearbySharingServiceImpl::RunOnNearbySharingServiceThreadDelayed( void NearbySharingServiceImpl::RunOnAnyThread(absl::string_view task_name, absl::AnyInvocable task) { - if (is_shutting_down_ == nullptr || *is_shutting_down_) { + if (IsShuttingDown()) { NL_LOG(WARNING) << __func__ << ": Skip the task " << task_name << " due to service is shutting down."; return; diff --git a/sharing/nearby_sharing_service_impl.h b/sharing/nearby_sharing_service_impl.h index e2073df3..a8d8634f 100644 --- a/sharing/nearby_sharing_service_impl.h +++ b/sharing/nearby_sharing_service_impl.h @@ -386,8 +386,7 @@ class NearbySharingServiceImpl const nearby::sharing::service::proto::ProgressUpdateFrame& progress_update_frame); - void OnConnectionDisconnected(int64_t share_target_id, - TransferMetadata::Status status); + void OnConnectionDisconnected(int64_t share_target_id); void OnIncomingMutualAcceptanceTimeout(int64_t share_target_id); void OnOutgoingMutualAcceptanceTimeout(int64_t share_target_id); @@ -406,8 +405,6 @@ class NearbySharingServiceImpl void RemoveIncomingPayloads(ShareTarget share_target); void Disconnect(int64_t share_target_id, TransferMetadata metadata); void OnDisconnectingConnectionTimeout(absl::string_view endpoint_id); - void OnDisconnectingConnectionDisconnected(int64_t share_target_id, - absl::string_view endpoint_id); ShareTargetInfo& GetOrCreateShareTargetInfo(const ShareTarget& share_target, absl::string_view endpoint_id); @@ -483,6 +480,8 @@ class NearbySharingServiceImpl // Update file path for the file attachment. void UpdateFilePath(ShareTarget& share_target); + // Returns true if Shutdown() has been called. + bool IsShuttingDown(); // Used to run nearby sharing service APIs. std::unique_ptr service_thread_; diff --git a/sharing/nearby_sharing_service_impl_test.cc b/sharing/nearby_sharing_service_impl_test.cc index a3cd6b7e..15306b5f 100644 --- a/sharing/nearby_sharing_service_impl_test.cc +++ b/sharing/nearby_sharing_service_impl_test.cc @@ -2354,6 +2354,62 @@ TEST_F(NearbySharingServiceImplTest, UnregisterReceiveSurfaceNeverRegistered) { EXPECT_FALSE(fake_nearby_connections_manager_->IsAdvertising()); } +TEST_F(NearbySharingServiceImplTest, + IncomingConnectionClosedAfterShutdown) { + fake_nearby_connections_manager_->SetRawAuthenticationToken(kEndpointId, + GetToken()); + SetUpAdvertisementDecoder(GetValidV1EndpointInfo(), + /*return_empty_advertisement=*/false, + /*return_empty_device_name=*/false, + /*expected_number_of_calls=*/1u); + + SetConnectionType(ConnectionType::kWifi); + NiceMock callback; + EXPECT_CALL(callback, OnTransferUpdate(testing::_, testing::_)).Times(0); + + SetUpForegroundReceiveSurface(callback); + EXPECT_CALL(*mock_app_info_, SetActiveFlag()); + Shutdown(); + + service_->OnIncomingConnection(kEndpointId, GetValidV1EndpointInfo(), + &connection_); + + sharing_service_task_runner_->SyncWithTimeout(kTaskWaitTimeout); + service_.reset(); +} + +TEST_F(NearbySharingServiceImplTest, + IncomingConnectionClosedBeforeCertDecryption) { + fake_nearby_connections_manager_->SetRawAuthenticationToken(kEndpointId, + GetToken()); + SetUpAdvertisementDecoder(GetValidV1EndpointInfo(), + /*return_empty_advertisement=*/false, + /*return_empty_device_name=*/false, + /*expected_number_of_calls=*/1u); + + SetConnectionType(ConnectionType::kWifi); + NiceMock callback; + EXPECT_CALL(callback, OnTransferUpdate(testing::_, testing::_)) + .WillOnce(testing::Invoke( + [](const ShareTarget& share_target, TransferMetadata metadata) { + EXPECT_TRUE(metadata.is_final_status()); + EXPECT_EQ(TransferMetadata::Status::kAwaitingRemoteAcceptanceFailed, + metadata.status()); + })); + + SetUpForegroundReceiveSurface(callback); + EXPECT_CALL(*mock_app_info_, SetActiveFlag()); + service_->OnIncomingConnection(kEndpointId, GetValidV1EndpointInfo(), + &connection_); + sharing_service_task_runner_->PostTask([this]() { + connection_.Close(); + }); + sharing_service_task_runner_->SyncWithTimeout(kTaskWaitTimeout); + + // To avoid UAF in OnIncomingTransferUpdate(). + UnregisterReceiveSurface(&callback); +} + TEST_F(NearbySharingServiceImplTest, IncomingConnectionClosedReadingIntroduction) { fake_nearby_connections_manager_->SetRawAuthenticationToken(kEndpointId, diff --git a/sharing/share_target_info.cc b/sharing/share_target_info.cc index 8503d42b..7276640b 100644 --- a/sharing/share_target_info.cc +++ b/sharing/share_target_info.cc @@ -21,6 +21,7 @@ #include "sharing/internal/public/logging.h" #include "sharing/share_target.h" #include "sharing/transfer_metadata.h" +#include "sharing/transfer_metadata_builder.h" namespace nearby { namespace sharing { @@ -32,7 +33,7 @@ ShareTargetInfo::ShareTargetInfo( : endpoint_id_(std::move(endpoint_id)), self_share_(share_target.for_self_share), share_target_(share_target), - transfer_update_callback_(std::move(transfer_update_callback)){} + transfer_update_callback_(std::move(transfer_update_callback)) {} ShareTargetInfo::ShareTargetInfo(ShareTargetInfo&&) = default; @@ -64,5 +65,23 @@ void ShareTargetInfo::UpdateTransferMetadata( } } +void ShareTargetInfo::set_disconnect_status( + TransferMetadata::Status disconnect_status) { + disconnect_status_ = disconnect_status; + if (disconnect_status_ != TransferMetadata::Status::kUnknown && + !TransferMetadata::IsFinalStatus(disconnect_status_)) { + NL_LOG(DFATAL) << "Disconnect status is not final: " + << static_cast(disconnect_status_); + } +} + +void ShareTargetInfo::OnDisconnect() { + if (disconnect_status_ != TransferMetadata::Status::kUnknown) { + UpdateTransferMetadata( + TransferMetadataBuilder().set_status(disconnect_status_).build()); + } + connection_ = nullptr; +} + } // namespace sharing } // namespace nearby diff --git a/sharing/share_target_info.h b/sharing/share_target_info.h index 1752f85d..22624ecc 100644 --- a/sharing/share_target_info.h +++ b/sharing/share_target_info.h @@ -120,6 +120,17 @@ class ShareTargetInfo { ShareTarget share_target() const { return share_target_; } + // Sets the status to send in the TransferMetadataUpdate on connection + // disconnect. If |status| is kUnknown, then no TransferMetadataUpdate will be + // sent. If |status| is set, it must be a final status. + void set_disconnect_status(TransferMetadata::Status disconnect_status); + + TransferMetadata::Status disconnect_status() const { + return disconnect_status_; + } + + void OnDisconnect(); + private: std::string endpoint_id_; std::optional certificate_; @@ -137,6 +148,10 @@ class ShareTargetInfo { bool got_final_status_ = false; std::function transfer_update_callback_; + // The status sent in the TransferMetadataUpdate on connection disconnect. + // If status is kUnknown, then no TransferMetadataUpdate will be sent. + TransferMetadata::Status disconnect_status_ = + TransferMetadata::Status::kUnknown; }; } // namespace sharing diff --git a/sharing/share_target_info_test.cc b/sharing/share_target_info_test.cc new file mode 100644 index 00000000..a2883453 --- /dev/null +++ b/sharing/share_target_info_test.cc @@ -0,0 +1,124 @@ +// Copyright 2024 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 "sharing/share_target_info.h" + +#include +#include +#include + +#include "gmock/gmock.h" +#include "protobuf-matchers/protocol-buffer-matchers.h" +#include "gtest/gtest.h" +#include "absl/strings/string_view.h" +#include "sharing/share_target.h" +#include "sharing/transfer_metadata.h" +#include "sharing/transfer_metadata_builder.h" + +namespace nearby::sharing { +namespace { +using testing::_; +using testing::Eq; +using testing::Invoke; +using testing::IsTrue; +using testing::MockFunction; + +constexpr absl::string_view kEndpointId = "12345"; + +class TestShareTargetInfo : public ShareTargetInfo { + public: + TestShareTargetInfo( + std::string endpoint_id, const ShareTarget& share_target, + std::function + transfer_update_callback) + : ShareTargetInfo(std::move(endpoint_id), share_target, + std::move(transfer_update_callback)), + is_incoming_(share_target.is_incoming) {} + + bool IsIncoming() const override { return is_incoming_; } + + private: + const bool is_incoming_; +}; + +TEST(ShareTargetInfoTest, UpdateTransferMetadata) { + MockFunction + update_callback; + ShareTarget share_target; + TestShareTargetInfo info(std::string(kEndpointId), share_target, + update_callback.AsStdFunction()); + + EXPECT_CALL(update_callback, Call(_, _)).Times(2); + + info.UpdateTransferMetadata( + TransferMetadataBuilder() + .set_status(TransferMetadata::Status::kInProgress) + .build()); + info.UpdateTransferMetadata( + TransferMetadataBuilder() + .set_status(TransferMetadata::Status::kInProgress) + .build()); +} + +TEST(ShareTargetInfoTest, UpdateTransferMetadataAfterFinalStatus) { + MockFunction + update_callback; + ShareTarget share_target; + TestShareTargetInfo info(std::string(kEndpointId), share_target, + update_callback.AsStdFunction()); + + EXPECT_CALL(update_callback, Call(_, _)); + + info.UpdateTransferMetadata( + TransferMetadataBuilder() + .set_status(TransferMetadata::Status::kComplete) + .build()); + info.UpdateTransferMetadata( + TransferMetadataBuilder() + .set_status(TransferMetadata::Status::kInProgress) + .build()); +} + +TEST(ShareTargetInfoTest, SetDisconnectStatus) { + MockFunction + update_callback; + ShareTarget share_target; + TestShareTargetInfo info(std::string(kEndpointId), share_target, + update_callback.AsStdFunction()); + + info.set_disconnect_status(TransferMetadata::Status::kCancelled); + EXPECT_EQ(info.disconnect_status(), TransferMetadata::Status::kCancelled); +} + +TEST(ShareTargetInfoTest, OnDisconnect) { + MockFunction + update_callback; + ShareTarget share_target; + TestShareTargetInfo info(std::string(kEndpointId), share_target, + update_callback.AsStdFunction()); + info.set_disconnect_status(TransferMetadata::Status::kCancelled); + EXPECT_EQ(info.disconnect_status(), TransferMetadata::Status::kCancelled); + EXPECT_CALL(update_callback, Call(_, _)) + .WillOnce(Invoke([](const ShareTarget& share_target, + const TransferMetadata& transfer_metadata) { + EXPECT_THAT(transfer_metadata.status(), + Eq(TransferMetadata::Status::kCancelled)); + EXPECT_THAT(transfer_metadata.is_final_status(), IsTrue()); + })); + + info.OnDisconnect(); +} + +} // namespace +} // namespace nearby::sharing