From a9d706179497ef702247b83b306166a80db673e1 Mon Sep 17 00:00:00 2001 From: hai007 Date: Mon, 13 Jun 2022 11:47:51 -0700 Subject: [PATCH] Group Connection Info into structure PiperOrigin-RevId: 454663252 --- connections/connection_options.h | 13 +++ .../implementation/base_pcp_handler.cc | 79 +++++++++---------- connections/implementation/base_pcp_handler.h | 17 ++-- .../implementation/endpoint_manager_test.cc | 17 +++- connections/implementation/offline_frames.cc | 45 +++++------ connections/implementation/offline_frames.h | 10 +-- .../implementation/offline_frames_test.cc | 19 +++-- .../offline_frames_validator_test.cc | 76 ++++++++---------- 8 files changed, 139 insertions(+), 137 deletions(-) diff --git a/connections/connection_options.h b/connections/connection_options.h index a94b1731..35f5b441 100644 --- a/connections/connection_options.h +++ b/connections/connection_options.h @@ -42,6 +42,19 @@ struct ConnectionOptions : public OptionsBase { std::vector GetMediums() const; }; +struct ConnectionInfo { + std::string local_endpoint_id; + ByteArray local_endpoint_info; + std::int32_t nonce; + bool supports_5_ghz = false; + std::string bssid; + std::int32_t ap_frequency = -1; + std::string ip_address; + std::vector supported_mediums; + std::int32_t keep_alive_interval_millis; + std::int32_t keep_alive_timeout_millis; +}; + } // namespace connections } // namespace nearby } // namespace location diff --git a/connections/implementation/base_pcp_handler.cc b/connections/implementation/base_pcp_handler.cc index 1b76b0fb..ee05449d 100644 --- a/connections/implementation/base_pcp_handler.cc +++ b/connections/implementation/base_pcp_handler.cc @@ -20,8 +20,8 @@ #include #include #include -#include #include +#include #include "securegcm/d2d_connection_context_v1.h" #include "securegcm/ukey2_handshake.h" @@ -457,6 +457,37 @@ void BasePcpHandler::OnEncryptionFailureRunnable( info.result.lock().get()); } +ConnectionInfo BasePcpHandler::FillConnectionInfo( + ClientProxy* client, const ConnectionRequestInfo& info, + const ConnectionOptions& connection_options) { + ConnectionInfo connection_info; + connection_info.local_endpoint_id = client->GetLocalEndpointId(); + connection_info.local_endpoint_info = info.endpoint_info; + connection_info.nonce = prng_.NextInt32(); + if (mediums_->GetWifi().IsAvailable()) { + connection_info.supports_5_ghz = + mediums_->GetWifi().GetCapability().supports_5_ghz; + + api::WifiInformation& wifi_info = mediums_->GetWifi().GetInformation(); + connection_info.bssid = wifi_info.bssid; + connection_info.ap_frequency = wifi_info.ap_frequency; + connection_info.ip_address = wifi_info.ip_address_4_bytes; + NEARBY_LOGS(INFO) << "Query for WIFI information: is_supports_5_ghz=" + << connection_info.supports_5_ghz + << "; bssid=" << connection_info.bssid + << "; ap_frequency=" << connection_info.ap_frequency + << "Mhz; ip_address in bytes format=" + << connection_info.ip_address; + } + connection_info.supported_mediums = + GetSupportedConnectionMediumsByPriority(connection_options); + connection_info.keep_alive_interval_millis = + connection_options.keep_alive_interval_millis; + connection_info.keep_alive_timeout_millis = + connection_options.keep_alive_timeout_millis; + return connection_info; +} + Status BasePcpHandler::RequestConnection( ClientProxy* client, const std::string& endpoint_id, const ConnectionRequestInfo& info, @@ -545,36 +576,13 @@ Status BasePcpHandler::RequestConnection( << "In requestConnection(), wrote ConnectionRequestFrame " "to endpoint_id=" << endpoint_id; - // Generate the nonce to use for this connection. - std::int32_t nonce = prng_.NextInt32(); - bool is_supports_5_ghz = false; - std::string bssid = ""; - std::int32_t ap_frequency = -1; - std::string ip_address = ""; - if (mediums_->GetWifi().IsAvailable()) { - is_supports_5_ghz = - mediums_->GetWifi().GetCapability().supports_5_ghz; + ConnectionInfo connection_info = + FillConnectionInfo(client, info, connection_options); - api::WifiInformation& wifi_info = - mediums_->GetWifi().GetInformation(); - bssid = wifi_info.bssid; - ap_frequency = wifi_info.ap_frequency; - ip_address = wifi_info.ip_address_4_bytes; - NEARBY_LOGS(INFO) << "Query for WIFI information: is_supports_5_ghz=" - << is_supports_5_ghz << "; bssid=" << bssid - << "; ap_frequency=" << ap_frequency - << "Mhz; ip_address in bytes format=" << ip_address; - } + Exception write_exception = + WriteConnectionRequestFrame(connection_info, channel.get()); - // The first message we have to send, after connecting, is to tell the - // endpoint about ourselves. - Exception write_exception = WriteConnectionRequestFrame( - channel.get(), client->GetLocalEndpointId(), info.endpoint_info, - nonce, is_supports_5_ghz, bssid, ap_frequency, ip_address, - GetSupportedConnectionMediumsByPriority(connection_options), - connection_options.keep_alive_interval_millis, - connection_options.keep_alive_timeout_millis); if (!write_exception.Ok()) { NEARBY_LOGS(INFO) << "Failed to send connection request: endpoint_id=" << endpoint_id; @@ -598,7 +606,7 @@ Status BasePcpHandler::RequestConnection( PendingConnectionInfo pendingConnectionInfo{}; pendingConnectionInfo.client = client; pendingConnectionInfo.remote_endpoint_info = endpoint->endpoint_info; - pendingConnectionInfo.nonce = nonce; + pendingConnectionInfo.nonce = connection_info.nonce; pendingConnectionInfo.is_incoming = false; pendingConnectionInfo.start_time = start_time; pendingConnectionInfo.listener = info.listener; @@ -728,17 +736,8 @@ bool BasePcpHandler::CanReceiveIncomingConnection(ClientProxy* client) const { } Exception BasePcpHandler::WriteConnectionRequestFrame( - EndpointChannel* endpoint_channel, const std::string& local_endpoint_id, - const ByteArray& local_endpoint_info, std::int32_t nonce, - bool supports_5_ghz, const std::string& bssid, std::int32_t ap_frequency, - const std::string& ip_address, - const std::vector& supported_mediums, - std::int32_t keep_alive_interval_millis, - std::int32_t keep_alive_timeout_millis) { - return endpoint_channel->Write(parser::ForConnectionRequest( - local_endpoint_id, local_endpoint_info, nonce, supports_5_ghz, bssid, - ap_frequency, ip_address, supported_mediums, keep_alive_interval_millis, - keep_alive_timeout_millis)); + const ConnectionInfo& conection_info, EndpointChannel* endpoint_channel) { + return endpoint_channel->Write(parser::ForConnectionRequest(conection_info)); } void BasePcpHandler::ProcessPreConnectionInitiationFailure( diff --git a/connections/implementation/base_pcp_handler.h b/connections/implementation/base_pcp_handler.h index a2023edc..517f74db 100644 --- a/connections/implementation/base_pcp_handler.h +++ b/connections/implementation/base_pcp_handler.h @@ -39,12 +39,12 @@ #include "connections/implementation/pcp_handler.h" #include "connections/listeners.h" #include "connections/status.h" -#include "internal/platform/byte_array.h" -#include "internal/platform/prng.h" #include "internal/platform/atomic_boolean.h" +#include "internal/platform/byte_array.h" #include "internal/platform/cancelable_alarm.h" #include "internal/platform/count_down_latch.h" #include "internal/platform/future.h" +#include "internal/platform/prng.h" #include "internal/platform/scheduled_executor.h" #include "internal/platform/single_thread_executor.h" @@ -108,6 +108,10 @@ class BasePcpHandler : public PcpHandler, void InjectEndpoint(ClientProxy* client, const std::string& service_id, const OutOfBandConnectionMetadata& metadata) override; + ConnectionInfo FillConnectionInfo( + ClientProxy* client, const ConnectionRequestInfo& info, + const ConnectionOptions& connection_options); + // Requests a newly discovered remote endpoint it to form a connection. // Updates state on ClientProxy. Status RequestConnection( @@ -386,14 +390,7 @@ class BasePcpHandler : public PcpHandler, EndpointChannel* endpoint_channel); static Exception WriteConnectionRequestFrame( - EndpointChannel* endpoint_channel, const std::string& local_endpoint_id, - const ByteArray& local_endpoint_info, std::int32_t nonce, - bool supports_5_ghz, const std::string& bssid, std::int32_t ap_frequency, - const std::string& ip_address, - const std::vector& supported_mediums, - std::int32_t keep_alive_interval_millis, - std::int32_t keep_alive_timeout_millis); - + const ConnectionInfo& conection_info, EndpointChannel* endpoint_channel); static constexpr absl::Duration kConnectionRequestReadTimeout = absl::Seconds(2); static constexpr absl::Duration kRejectedConnectionCloseDelay = diff --git a/connections/implementation/endpoint_manager_test.cc b/connections/implementation/endpoint_manager_test.cc index 7c827006..0b28c8ad 100644 --- a/connections/implementation/endpoint_manager_test.cc +++ b/connections/implementation/endpoint_manager_test.cc @@ -195,10 +195,19 @@ TEST_F(EndpointManagerTest, RegisterFrameProcessorWorks) { auto endpoint_channel = std::make_unique(); auto connect_request = std::make_unique(); ByteArray endpoint_info{"endpoint_name"}; - auto read_data = parser::ForConnectionRequest( - "endpoint_id", endpoint_info, 1234, false, "", 2412, "8xqT", - std::vector{Medium::BLE}, 0, - 0); + ConnectionInfo connection_info{ + "endpoint_id", + endpoint_info, + 1234 /*nonce*/, + false /*supports_5_ghz*/, + "" /*bssid*/, + 2412 /*ap_frequency*/, + "8xqT" /*ip_address in 4 bytes format*/, + std::vector{Medium::BLE} /*supported_mediums*/, + 0 /*keep_alive_interval_millis*/, + 0 /*keep_alive_timeout_millis*/}; + + auto read_data = parser::ForConnectionRequest(connection_info); EXPECT_CALL(*connect_request, OnIncomingFrame); EXPECT_CALL(*connect_request, OnEndpointDisconnect); EXPECT_CALL(*endpoint_channel, Read()) diff --git a/connections/implementation/offline_frames.cc b/connections/implementation/offline_frames.cc index 5294a482..2e7f808c 100644 --- a/connections/implementation/offline_frames.cc +++ b/connections/implementation/offline_frames.cc @@ -62,44 +62,41 @@ V1Frame::FrameType GetFrameType(const OfflineFrame& frame) { return V1Frame::UNKNOWN_FRAME_TYPE; } -ByteArray ForConnectionRequest(const std::string& endpoint_id, - const ByteArray& endpoint_info, - std::int32_t nonce, bool supports_5_ghz, - const std::string& bssid, - std::int32_t ap_frequency, - const std::string& ip_address, - const std::vector& mediums, - std::int32_t keep_alive_interval_millis, - std::int32_t keep_alive_timeout_millis) { +ByteArray ForConnectionRequest(const ConnectionInfo& conection_info) { OfflineFrame frame; frame.set_version(OfflineFrame::V1); auto* v1_frame = frame.mutable_v1(); v1_frame->set_type(V1Frame::CONNECTION_REQUEST); auto* connection_request = v1_frame->mutable_connection_request(); - if (!endpoint_id.empty()) connection_request->set_endpoint_id(endpoint_id); - if (!endpoint_info.Empty()) { - connection_request->set_endpoint_name(std::string(endpoint_info)); - connection_request->set_endpoint_info(std::string(endpoint_info)); + if (!conection_info.local_endpoint_id.empty()) + connection_request->set_endpoint_id(conection_info.local_endpoint_id); + if (!conection_info.local_endpoint_info.Empty()) { + connection_request->set_endpoint_name( + std::string(conection_info.local_endpoint_info)); + connection_request->set_endpoint_info( + std::string(conection_info.local_endpoint_info)); } - connection_request->set_nonce(nonce); + connection_request->set_nonce(conection_info.nonce); auto* medium_metadata = connection_request->mutable_medium_metadata(); - medium_metadata->set_supports_5_ghz(supports_5_ghz); - if (!bssid.empty()) medium_metadata->set_bssid(bssid); - medium_metadata->set_ap_frequency(ap_frequency); - if (!ip_address.empty()) medium_metadata->set_ip_address(ip_address); - if (!mediums.empty()) { - for (const auto& medium : mediums) { + medium_metadata->set_supports_5_ghz(conection_info.supports_5_ghz); + if (!conection_info.bssid.empty()) + medium_metadata->set_bssid(conection_info.bssid); + medium_metadata->set_ap_frequency(conection_info.ap_frequency); + if (!conection_info.ip_address.empty()) + medium_metadata->set_ip_address(conection_info.ip_address); + if (!conection_info.supported_mediums.empty()) { + for (const auto& medium : conection_info.supported_mediums) { connection_request->add_mediums(MediumToConnectionRequestMedium(medium)); } } - if (keep_alive_interval_millis > 0) { + if (conection_info.keep_alive_interval_millis > 0) { connection_request->set_keep_alive_interval_millis( - keep_alive_interval_millis); + conection_info.keep_alive_interval_millis); } - if (keep_alive_timeout_millis > 0) { + if (conection_info.keep_alive_timeout_millis > 0) { connection_request->set_keep_alive_timeout_millis( - keep_alive_timeout_millis); + conection_info.keep_alive_timeout_millis); } return ToBytes(std::move(frame)); diff --git a/connections/implementation/offline_frames.h b/connections/implementation/offline_frames.h index 462de6f4..ef3c8b46 100644 --- a/connections/implementation/offline_frames.h +++ b/connections/implementation/offline_frames.h @@ -42,15 +42,7 @@ ExceptionOr FromBytes(const ByteArray& offline_frame_bytes); V1Frame::FrameType GetFrameType(const OfflineFrame& offline_frame); // Builds Connection Request / Response messages. -ByteArray ForConnectionRequest(const std::string& endpoint_id, - const ByteArray& endpoint_info, - std::int32_t nonce, bool supports_5_ghz, - const std::string& bssid, - std::int32_t ap_frequency, - const std::string& ip_address, - const std::vector& mediums, - std::int32_t keep_alive_interval_millis, - std::int32_t keep_alive_timeout_millis); +ByteArray ForConnectionRequest(const ConnectionInfo& conection_info); ByteArray ForConnectionResponse(std::int32_t status); // Builds Payload transfer messages. diff --git a/connections/implementation/offline_frames_test.cc b/connections/implementation/offline_frames_test.cc index d434b6cf..da17d2ce 100644 --- a/connections/implementation/offline_frames_test.cc +++ b/connections/implementation/offline_frames_test.cc @@ -111,12 +111,19 @@ TEST(OfflineFramesTest, CanGenerateConnectionRequest) { keep_alive_timeout_millis: 5000 > >)pb"; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + + ConnectionInfo connection_info{std::string(kEndpointId), + ByteArray{std::string(kEndpointName)}, + kNonce, + kSupports5ghz, + std::string(kBssid), + kApFrequency, + std::string(kIp4Bytes), + std::vector>( + kMediums.begin(), kMediums.end()), + kKeepAliveIntervalMillis, + kKeepAliveTimeoutMillis}; + ByteArray bytes = ForConnectionRequest(connection_info); auto response = FromBytes(bytes); ASSERT_TRUE(response.ok()); OfflineFrame message = FromBytes(bytes).result(); diff --git a/connections/implementation/offline_frames_validator_test.cc b/connections/implementation/offline_frames_validator_test.cc index 04cd73e6..ffae3d8d 100644 --- a/connections/implementation/offline_frames_validator_test.cc +++ b/connections/implementation/offline_frames_validator_test.cc @@ -54,15 +54,26 @@ constexpr std::array kMediums = { constexpr int kKeepAliveIntervalMillis = 1000; constexpr int kKeepAliveTimeoutMillis = 5000; -TEST(OfflineFramesValidatorTest, ValidatesAsOkWithValidConnectionRequestFrame) { +class OfflineFramesConnectionRequestTest : public testing::Test { + protected: + ConnectionInfo connection_info_{std::string(kEndpointId), + ByteArray{std::string(kEndpointName)}, + kNonce, + kSupports5ghz, + std::string(kBssid), + kApFrequency, + std::string(kIp4Bytes), + std::vector>( + kMediums.begin(), kMediums.end()), + kKeepAliveIntervalMillis, + kKeepAliveTimeoutMillis}; +}; + +TEST_F(OfflineFramesConnectionRequestTest, + ValidatesAsOkWithValidConnectionRequestFrame) { OfflineFrame offline_frame; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto ret_value = EnsureValidOfflineFrame(offline_frame); @@ -70,16 +81,11 @@ TEST(OfflineFramesValidatorTest, ValidatesAsOkWithValidConnectionRequestFrame) { ASSERT_TRUE(ret_value.Ok()); } -TEST(OfflineFramesValidatorTest, +TEST_F(OfflineFramesConnectionRequestTest, ValidatesAsFailWithNullConnectionRequestFrame) { OfflineFrame offline_frame; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto* v1_frame = offline_frame.mutable_v1(); @@ -90,17 +96,12 @@ TEST(OfflineFramesValidatorTest, ASSERT_FALSE(ret_value.Ok()); } -TEST(OfflineFramesValidatorTest, +TEST_F(OfflineFramesConnectionRequestTest, ValidatesAsFailWithNullEndpointIdInConnectionRequestFrame) { OfflineFrame offline_frame; - std::string empty_enpoint_id; - ByteArray bytes = ForConnectionRequest( - empty_enpoint_id, ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + connection_info_.local_endpoint_id = ""; + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto ret_value = EnsureValidOfflineFrame(offline_frame); @@ -108,17 +109,12 @@ TEST(OfflineFramesValidatorTest, ASSERT_FALSE(ret_value.Ok()); } -TEST(OfflineFramesValidatorTest, +TEST_F(OfflineFramesConnectionRequestTest, ValidatesAsFailWithNullEndpointInfoInConnectionRequestFrame) { OfflineFrame offline_frame; - ByteArray empty_endpoint_info; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), empty_endpoint_info, kNonce, kSupports5ghz, - std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + connection_info_.local_endpoint_info = ByteArray{""}; + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto ret_value = EnsureValidOfflineFrame(offline_frame); @@ -126,17 +122,12 @@ TEST(OfflineFramesValidatorTest, ASSERT_FALSE(ret_value.Ok()); } -TEST(OfflineFramesValidatorTest, +TEST_F(OfflineFramesConnectionRequestTest, ValidatesAsOkWithNullBssidInConnectionRequestFrame) { OfflineFrame offline_frame; - std::string empty_bssid; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, empty_bssid, kApFrequency, std::string(kIp4Bytes), - std::vector>(kMediums.begin(), - kMediums.end()), - kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + connection_info_.bssid = ""; + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto ret_value = EnsureValidOfflineFrame(offline_frame); @@ -144,15 +135,12 @@ TEST(OfflineFramesValidatorTest, ASSERT_TRUE(ret_value.Ok()); } -TEST(OfflineFramesValidatorTest, +TEST_F(OfflineFramesConnectionRequestTest, ValidatesAsOkWithNullMediumsInConnectionRequestFrame) { OfflineFrame offline_frame; - std::vector empty_mediums; - ByteArray bytes = ForConnectionRequest( - std::string(kEndpointId), ByteArray{std::string(kEndpointName)}, kNonce, - kSupports5ghz, std::string(kBssid), kApFrequency, std::string(kIp4Bytes), - empty_mediums, kKeepAliveIntervalMillis, kKeepAliveTimeoutMillis); + connection_info_.supported_mediums = {}; + ByteArray bytes = ForConnectionRequest(connection_info_); offline_frame.ParseFromString(std::string(bytes)); auto ret_value = EnsureValidOfflineFrame(offline_frame);