Move accept transfer into IncomingShareSession.

PiperOrigin-RevId: 651226239
This commit is contained in:
Francis Tsui
2024-07-10 19:09:04 -07:00
committed by Copybara-Service
parent 5d952cff5e
commit 76651429b8
5 changed files with 136 additions and 56 deletions
+31 -1
View File
@@ -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<void(int64_t, TransferMetadata)> update_callback) {
const absl::flat_hash_map<int64_t, int64_t>& 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<std::filesystem::path> 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
+8 -1
View File
@@ -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<void(int64_t, TransferMetadata)> update_callback);
@@ -82,6 +83,10 @@ class IncomingShareSession : public ShareSession {
// Returns the file paths of all file payloads.
std::vector<std::filesystem::path> 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<void(const IncomingShareSession&, const TransferMetadata&)>
transfer_update_callback_;
bool bandwidth_upgrade_requested_ = false;
};
} // namespace nearby::sharing
+72 -5
View File
@@ -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<void(const IncomingShareSession&, const TransferMetadata&)>
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<uint8_t> 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
+25 -46
View File
@@ -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<void(StatusCodes status_codes)> 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();
}
-3
View File
@@ -312,9 +312,6 @@ class NearbySharingServiceImpl
void OnTransferComplete();
void OnTransferStarted(bool is_incoming);
void ReceivePayloads(
IncomingShareSession& session,
std::function<void(StatusCodes status_codes)> status_codes_callback);
StatusCodes SendPayloads(ShareSession& session);
void OnOutgoingConnection(absl::Time connect_start_time,