diff --git a/connections/implementation/mediums/ble.cc b/connections/implementation/mediums/ble.cc index a9d79609..739d680f 100644 --- a/connections/implementation/mediums/ble.cc +++ b/connections/implementation/mediums/ble.cc @@ -423,12 +423,15 @@ BleSocket Ble::Connect(BlePeripheral& peripheral, const std::string& service_id, ByteArray Ble::UnwrapAdvertisementBytes( const ByteArray& medium_advertisement_data) { - mediums::BleAdvertisement medium_ble_advertisement{medium_advertisement_data}; - if (!medium_ble_advertisement.IsValid()) { - return ByteArray{}; + auto medium_ble_advertisement_status_or = + mediums::BleAdvertisement::CreateBleAdvertisement( + medium_advertisement_data); + if (!medium_ble_advertisement_status_or.ok()) { + NEARBY_LOGS(INFO) << medium_ble_advertisement_status_or.status().ToString(); + return ByteArray(); } - return medium_ble_advertisement.GetData(); + return medium_ble_advertisement_status_or.value().GetData(); } } // namespace connections diff --git a/connections/implementation/mediums/ble_v2/ble_advertisement.cc b/connections/implementation/mediums/ble_v2/ble_advertisement.cc index 3abb8a79..d4b81b76 100644 --- a/connections/implementation/mediums/ble_v2/ble_advertisement.cc +++ b/connections/implementation/mediums/ble_v2/ble_advertisement.cc @@ -20,9 +20,12 @@ #include #include +#include "absl/status/status.h" +#include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "connections/implementation/mediums/ble_v2/ble_advertisement_header.h" #include "internal/platform/base_input_stream.h" +#include "internal/platform/byte_array.h" #include "internal/platform/logging.h" namespace nearby { @@ -81,92 +84,85 @@ void BleAdvertisement::DoInitialize(bool fast_advertisement, Version version, psm_ = psm; } -BleAdvertisement::BleAdvertisement(const ByteArray &ble_advertisement_bytes) { +absl::StatusOr BleAdvertisement::CreateBleAdvertisement( + const ByteArray &ble_advertisement_bytes) { if (ble_advertisement_bytes.Empty()) { - NEARBY_LOG(INFO, - "Cannot deserialize BleAdvertisement: null bytes passed in."); - return; + return absl::InvalidArgumentError( + "Cannot deserialize BleAdvertisement: null bytes passed in."); } if (ble_advertisement_bytes.size() < kVersionLength) { - NEARBY_LOG( - INFO, - "Cannot deserialize BleAdvertisement: expecting min %d raw bytes to " - "parse the version, got %" PRIu64, - kVersionLength, ble_advertisement_bytes.size()); - return; + return absl::InvalidArgumentError(absl::StrCat( + "Cannot deserialize BleAdvertisement: expecting min ", kVersionLength, + " bytes, got ", ble_advertisement_bytes.size())); } - ByteArray advertisement_bytes{ble_advertisement_bytes}; - BaseInputStream base_input_stream{advertisement_bytes}; + ByteArray advertisement_bytes(ble_advertisement_bytes); + BaseInputStream base_input_stream(advertisement_bytes); // The first 1 byte is supposed to be the version, socket version and the fast // advertisement flag. auto version_byte = static_cast(base_input_stream.ReadUint8()); - // Version. - version_ = static_cast((version_byte & kVersionBitmask) >> 5); - if (!IsSupportedVersion(version_)) { - NEARBY_LOG(INFO, - "Cannot deserialize BleAdvertisement: unsupported Version %u", - version_); - return; + Version version = static_cast((version_byte & kVersionBitmask) >> 5); + if (!IsSupportedVersion(version)) { + return absl::InvalidArgumentError(absl::StrCat( + "Cannot deserialize BleAdvertisement: unsupported Version ", version)); } - // Socket version. - socket_version_ = + SocketVersion socket_version = static_cast((version_byte & kSocketVersionBitmask) >> 2); - if (!IsSupportedSocketVersion(socket_version_)) { - NEARBY_LOG( - INFO, - "Cannot deserialize BleAdvertisement: unsupported SocketVersion %u", - socket_version_); - version_ = Version::kUndefined; - return; + if (!IsSupportedSocketVersion(socket_version)) { + return absl::InvalidArgumentError(absl::StrCat( + "Cannot deserialize BleAdvertisement: unsupported SocketVersion ", + socket_version)); } - // Fast advertisement flag. - fast_advertisement_ = + bool fast_advertisement = static_cast((version_byte & kFastAdvertisementFlagBitmask) >> 1); // The next 3 bytes are supposed to be the service_id_hash if not fast // advertisement. - if (!fast_advertisement_) { - service_id_hash_ = base_input_stream.ReadBytes(kServiceIdHashLength); + ByteArray service_id_hash; + if (!fast_advertisement) { + service_id_hash = base_input_stream.ReadBytes(kServiceIdHashLength); } // Data length. int expected_data_size = - fast_advertisement_ + fast_advertisement ? static_cast( base_input_stream.ReadBytes(kFastDataSizeLength).data()[0]) : static_cast(base_input_stream.ReadUint32()); if (expected_data_size < 0) { - NEARBY_LOG(INFO, - "Cannot deserialize BleAdvertisement: negative data size %d", - expected_data_size); - version_ = Version::kUndefined; - return; + return absl::InvalidArgumentError( + absl::StrCat("Cannot deserialize BleAdvertisement: negative data size ", + expected_data_size)); } // Data. // Check that the stated data size is the same as what we received. - data_ = base_input_stream.ReadBytes(expected_data_size); - if (data_.size() != expected_data_size) { - NEARBY_LOG(INFO, - "Cannot deserialize BleAdvertisement: expected data to be %u " - "bytes, got %" PRIu64 " bytes ", - expected_data_size, data_.size()); - version_ = Version::kUndefined; - return; + auto data = base_input_stream.ReadBytes(expected_data_size); + if (data.size() != expected_data_size) { + return absl::InvalidArgumentError(absl::StrCat( + "Cannot deserialize BleAdvertisement: expected data to be ", + expected_data_size, " bytes, got ", data.size())); } + BleAdvertisement ble_advertisement; + ble_advertisement.version_ = version; + ble_advertisement.socket_version_ = socket_version; + ble_advertisement.fast_advertisement_ = fast_advertisement; + ble_advertisement.service_id_hash_ = service_id_hash; + ble_advertisement.data_ = data; + // Device token. If the number of remaining bytes are valid for device token, // then read it. if (base_input_stream.IsAvailable(kDeviceTokenLength)) { - device_token_ = base_input_stream.ReadBytes(kDeviceTokenLength); + ble_advertisement.device_token_ = + base_input_stream.ReadBytes(kDeviceTokenLength); } else { // No device token no more optional field. - return; + return ble_advertisement; } // Extra fields, for backward compatible reason, put this field in the end of @@ -178,8 +174,9 @@ BleAdvertisement::BleAdvertisement(const ByteArray &ble_advertisement_bytes) { if (base_input_stream.IsAvailable(extra_fields_byte_number)) { BleExtraFields extra_fields{ base_input_stream.ReadBytes(extra_fields_byte_number)}; - psm_ = extra_fields.GetPsm(); + ble_advertisement.psm_ = extra_fields.GetPsm(); } + return ble_advertisement; } BleAdvertisement::operator ByteArray() const { @@ -243,12 +240,11 @@ bool BleAdvertisement::operator==(const BleAdvertisement &rhs) const { this->GetPsm() == rhs.GetPsm(); } -bool BleAdvertisement::IsSupportedVersion(Version version) const { +bool BleAdvertisement::IsSupportedVersion(Version version) { return version >= Version::kV1 && version <= Version::kV2; } -bool BleAdvertisement::IsSupportedSocketVersion( - SocketVersion socket_version) const { +bool BleAdvertisement::IsSupportedSocketVersion(SocketVersion socket_version) { return socket_version >= SocketVersion::kV1 && socket_version <= SocketVersion::kV2; } diff --git a/connections/implementation/mediums/ble_v2/ble_advertisement.h b/connections/implementation/mediums/ble_v2/ble_advertisement.h index 94a20aa6..5a470584 100644 --- a/connections/implementation/mediums/ble_v2/ble_advertisement.h +++ b/connections/implementation/mediums/ble_v2/ble_advertisement.h @@ -17,6 +17,7 @@ #include +#include "absl/status/statusor.h" #include "connections/implementation/mediums/ble_v2/ble_advertisement_header.h" #include "internal/platform/byte_array.h" @@ -73,7 +74,9 @@ class BleAdvertisement { const ByteArray &service_id_hash, const ByteArray &data, const ByteArray &device_token, int psm = BleAdvertisementHeader::kDefaultPsmValue); - explicit BleAdvertisement(const ByteArray &ble_advertisement_bytes); + + static absl::StatusOr CreateBleAdvertisement( + const ByteArray &ble_advertisement_bytes); BleAdvertisement(const BleAdvertisement &) = default; BleAdvertisement &operator=(const BleAdvertisement &) = default; BleAdvertisement(BleAdvertisement &&) = default; @@ -128,8 +131,8 @@ class BleAdvertisement { SocketVersion socket_version, const ByteArray &service_id_hash, const ByteArray &data, const ByteArray &device_token, int psm); - bool IsSupportedVersion(Version version) const; - bool IsSupportedSocketVersion(SocketVersion socket_version) const; + static bool IsSupportedVersion(Version version); + static bool IsSupportedSocketVersion(SocketVersion socket_version); void SerializeDataSize(bool fast_advertisement, char *data_size_bytes_write_ptr, size_t data_size) const; diff --git a/connections/implementation/mediums/ble_v2/ble_advertisement_test.cc b/connections/implementation/mediums/ble_v2/ble_advertisement_test.cc index 99131eb6..4ef352d6 100644 --- a/connections/implementation/mediums/ble_v2/ble_advertisement_test.cc +++ b/connections/implementation/mediums/ble_v2/ble_advertisement_test.cc @@ -17,8 +17,13 @@ #include #include +#include "gmock/gmock.h" +#include "protobuf-matchers/protocol-buffer-matchers.h" #include "gtest/gtest.h" #include "absl/hash/hash_testing.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "internal/platform/byte_array.h" namespace nearby { namespace connections { @@ -40,6 +45,9 @@ constexpr size_t kAdvertisementLength = 77; constexpr size_t kFastAdvertisementLength = 16; constexpr size_t kLongAdvertisementLength = kAdvertisementLength + 1000; +using ::absl::StatusCode::kInvalidArgument; +using ::testing::status::StatusIs; + TEST(BleAdvertisementTest, ConstructionWorksV1) { ByteArray service_id_hash{std::string(kServiceIDHashBytes)}; ByteArray data{std::string(kData)}; @@ -224,7 +232,11 @@ TEST(BleAdvertisementTest, ConstructionFromSerializedBytesWorks) { kVersion, kSocketVersion, service_id_hash, data, device_token}; ByteArray ble_advertisement_bytes{original_ble_advertisement}; - BleAdvertisement ble_advertisement{ble_advertisement_bytes}; + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); EXPECT_TRUE(ble_advertisement.IsValid()); EXPECT_FALSE(ble_advertisement.IsFastAdvertisement()); @@ -245,7 +257,11 @@ TEST(BleAdvertisementTest, kVersion, kSocketVersion, ByteArray{}, fast_data, device_token}; ByteArray ble_advertisement_bytes{original_ble_advertisement}; - BleAdvertisement ble_advertisement{ble_advertisement_bytes}; + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); EXPECT_TRUE(ble_advertisement.IsValid()); EXPECT_TRUE(ble_advertisement.IsFastAdvertisement()); @@ -263,7 +279,11 @@ TEST(BleAdvertisementTest, ConstructionFromSerializedBytesWithEmptyDataWorks) { BleAdvertisement original_ble_advertisement{ kVersion, kSocketVersion, service_id_hash, ByteArray(), device_token}; ByteArray ble_advertisement_bytes{original_ble_advertisement}; - BleAdvertisement ble_advertisement{ble_advertisement_bytes}; + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); EXPECT_TRUE(ble_advertisement.IsValid()); EXPECT_FALSE(ble_advertisement.IsFastAdvertisement()); @@ -281,7 +301,11 @@ TEST(BleAdvertisementTest, BleAdvertisement original_ble_advertisement{ kVersion, kSocketVersion, ByteArray{}, ByteArray(), device_token}; ByteArray ble_advertisement_bytes{original_ble_advertisement}; - BleAdvertisement ble_advertisement{ble_advertisement_bytes}; + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); EXPECT_TRUE(ble_advertisement.IsValid()); EXPECT_TRUE(ble_advertisement.IsFastAdvertisement()); @@ -310,7 +334,11 @@ TEST(BleAdvertisementTest, ConstructionFromExtraSerializedBytesWorks) { // Re-parse the Ble advertisement using our extra long advertisement bytes. ByteArray long_ble_advertisement_bytes{raw_ble_advertisement_bytes, kLongAdvertisementLength}; - BleAdvertisement long_ble_advertisement{long_ble_advertisement_bytes}; + + auto long_ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(long_ble_advertisement_bytes); + ASSERT_OK(long_ble_advertisement_status_or.status()); + auto long_ble_advertisement = long_ble_advertisement_status_or.value(); EXPECT_TRUE(long_ble_advertisement.IsValid()); EXPECT_FALSE(long_ble_advertisement.IsFastAdvertisement()); @@ -341,7 +369,11 @@ TEST(BleAdvertisementTest, // Re-parse the Ble advertisement using our extra long advertisement bytes. ByteArray long_ble_advertisement_bytes{raw_ble_advertisement_bytes, kLongAdvertisementLength}; - BleAdvertisement long_ble_advertisement{long_ble_advertisement_bytes}; + + auto long_ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(long_ble_advertisement_bytes); + ASSERT_OK(long_ble_advertisement_status_or.status()); + auto long_ble_advertisement = long_ble_advertisement_status_or.value(); EXPECT_TRUE(long_ble_advertisement.IsValid()); EXPECT_TRUE(long_ble_advertisement.IsFastAdvertisement()); @@ -353,9 +385,8 @@ TEST(BleAdvertisementTest, } TEST(BleAdvertisementTest, ConstructionFromNullBytesFails) { - BleAdvertisement ble_advertisement{ByteArray{}}; - - EXPECT_FALSE(ble_advertisement.IsValid()); + EXPECT_THAT(BleAdvertisement::CreateBleAdvertisement(ByteArray()), + StatusIs(kInvalidArgument)); } TEST(BleAdvertisementTest, ConstructionFromShortLengthSerializedBytesFails) { @@ -370,9 +401,9 @@ TEST(BleAdvertisementTest, ConstructionFromShortLengthSerializedBytesFails) { // Cut off the advertisement so that it's too short. ByteArray short_ble_advertisement_bytes{ original_ble_advertisement_bytes.data(), 7}; - BleAdvertisement short_ble_advertisement{short_ble_advertisement_bytes}; - - EXPECT_FALSE(short_ble_advertisement.IsValid()); + EXPECT_THAT( + BleAdvertisement::CreateBleAdvertisement(short_ble_advertisement_bytes), + StatusIs(kInvalidArgument)); } TEST(BleAdvertisementTest, @@ -387,9 +418,9 @@ TEST(BleAdvertisementTest, // Cut off the advertisement so that it's too short. ByteArray short_ble_advertisement_bytes{ original_ble_advertisement_bytes.data(), 2}; - BleAdvertisement short_ble_advertisement{short_ble_advertisement_bytes}; - - EXPECT_FALSE(short_ble_advertisement.IsValid()); + EXPECT_THAT( + BleAdvertisement::CreateBleAdvertisement(short_ble_advertisement_bytes), + StatusIs(kInvalidArgument)); } TEST(BleAdvertisementTest, @@ -415,10 +446,9 @@ TEST(BleAdvertisementTest, // Try to parse the Ble advertisement using our corrupted advertisement bytes. ByteArray corrupted_ble_advertisement_bytes{raw_ble_advertisement_bytes, kAdvertisementLength}; - BleAdvertisement corrupted_ble_advertisement{ - corrupted_ble_advertisement_bytes}; - - EXPECT_FALSE(corrupted_ble_advertisement.IsValid()); + EXPECT_THAT(BleAdvertisement::CreateBleAdvertisement( + corrupted_ble_advertisement_bytes), + StatusIs(kInvalidArgument)); } TEST(BleAdvertisementTest, @@ -443,10 +473,9 @@ TEST(BleAdvertisementTest, // Try to parse the Ble advertisement using our corrupted advertisement bytes. ByteArray corrupted_ble_advertisement_bytes{raw_ble_advertisement_bytes, kFastAdvertisementLength}; - BleAdvertisement corrupted_ble_advertisement{ - corrupted_ble_advertisement_bytes}; - - EXPECT_FALSE(corrupted_ble_advertisement.IsValid()); + EXPECT_THAT(BleAdvertisement::CreateBleAdvertisement( + corrupted_ble_advertisement_bytes), + StatusIs(kInvalidArgument)); } TEST(BleAdvertisementTest, ConstructionWorksWithPsmValue) { @@ -480,7 +509,11 @@ TEST(BleAdvertisementTest, kVersion, kSocketVersion, service_id_hash, data, device_token, psm); ByteArray ble_advertisement_bytes = original_ble_advertisement.ByteArrayWithExtraField(); - BleAdvertisement ble_advertisement(ble_advertisement_bytes); + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); ASSERT_TRUE(ble_advertisement.IsValid()); EXPECT_FALSE(ble_advertisement.IsFastAdvertisement()); @@ -505,7 +538,11 @@ TEST(BleAdvertisementTest, // But use ByteArrayWithExtraField to restore back. It should fail. ByteArray ble_advertisement_bytes = original_ble_advertisement.ByteArrayWithExtraField(); - BleAdvertisement ble_advertisement(ble_advertisement_bytes); + + auto ble_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(ble_advertisement_bytes); + ASSERT_OK(ble_advertisement_status_or.status()); + auto ble_advertisement = ble_advertisement_status_or.value(); ASSERT_TRUE(ble_advertisement.IsValid()); EXPECT_FALSE(ble_advertisement.IsFastAdvertisement()); diff --git a/connections/implementation/mediums/ble_v2/discovered_peripheral_tracker.cc b/connections/implementation/mediums/ble_v2/discovered_peripheral_tracker.cc index b5b5f566..ba2f2a69 100644 --- a/connections/implementation/mediums/ble_v2/discovered_peripheral_tracker.cc +++ b/connections/implementation/mediums/ble_v2/discovered_peripheral_tracker.cc @@ -441,13 +441,14 @@ DiscoveredPeripheralTracker::ParseRawGattAdvertisements( // TODO(edwinwu): Refactor this big loop as subroutines. for (const auto gatt_advertisement_bytes : gatt_advertisement_bytes_list) { // First, parse the raw bytes into a BleAdvertisement. - BleAdvertisement gatt_advertisement(*gatt_advertisement_bytes); - if (!gatt_advertisement.IsValid()) { - NEARBY_LOGS(INFO) << "Unable to parse raw GATT advertisement:" - << absl::BytesToHexString( - gatt_advertisement_bytes->data()); + + auto gatt_advertisement_status_or = + BleAdvertisement::CreateBleAdvertisement(*gatt_advertisement_bytes); + if (!gatt_advertisement_status_or.ok()) { + NEARBY_LOGS(INFO) << gatt_advertisement_status_or.status().ToString(); continue; } + auto gatt_advertisement = gatt_advertisement_status_or.value(); // Make sure the advertisement belongs to a service ID we're tracking. for (const auto& item : service_id_infos_) {