diff --git a/connections/implementation/mediums/advertisements/BUILD b/connections/implementation/mediums/advertisements/BUILD index 70555806..09928b43 100644 --- a/connections/implementation/mediums/advertisements/BUILD +++ b/connections/implementation/mediums/advertisements/BUILD @@ -12,12 +12,17 @@ # See the License for the specific language governing permissions and # limitations under the License. +package(default_visibility = [ + "//:__subpackages__", +]) + licenses(["notice"]) cc_library( name = "common", srcs = ["data_element.cc"], hdrs = ["data_element.h"], + compatible_with = ["//buildenv/target:non_prod"], deps = [ "//internal/platform:base", "//internal/platform:logging", @@ -30,13 +35,14 @@ cc_library( name = "dct_advertisement", srcs = ["dct_advertisement.cc"], hdrs = ["dct_advertisement.h"], + compatible_with = ["//buildenv/target:non_prod"], deps = [ ":common", "//internal/crypto_cros", "//internal/platform:base", "//internal/platform:logging", + "//internal/platform:types", "//internal/platform:util", - "@com_google_absl//absl/random", "@com_google_absl//absl/strings", ], ) @@ -45,6 +51,7 @@ cc_library( name = "util", srcs = ["advertisement_util.cc"], hdrs = ["advertisement_util.h"], + compatible_with = ["//buildenv/target:non_prod"], deps = [ ":dct_advertisement", "//internal/platform:base", @@ -72,6 +79,7 @@ cc_test( ":dct_advertisement", "//internal/platform:base", "//internal/platform:util", + "//internal/platform/implementation/g3", "@com_github_protobuf_matchers//protobuf-matchers", "@com_google_googletest//:gtest_main", ], @@ -84,6 +92,7 @@ cc_test( ":util", "//internal/platform:base", "//internal/platform:util", + "//internal/platform/implementation/g3", "@com_github_protobuf_matchers//protobuf-matchers", "@com_google_googletest//:gtest_main", ], diff --git a/connections/implementation/mediums/advertisements/advertisement_util.cc b/connections/implementation/mediums/advertisements/advertisement_util.cc index c2bf0066..8366ac28 100644 --- a/connections/implementation/mediums/advertisements/advertisement_util.cc +++ b/connections/implementation/mediums/advertisements/advertisement_util.cc @@ -28,12 +28,12 @@ std::optional ReadDeviceName(const ByteArray& endpoint_info) { // LINT.IfChange StreamReader reader(endpoint_info); std::optional version = reader.ReadBits(3); - if (!version.has_value() || *version != 1) { + if (!version.has_value() || *version > 1) { return std::nullopt; } std::optional has_device_name = reader.ReadBits(1); - if (!has_device_name.has_value() || *has_device_name == 0) { + if (!has_device_name.has_value() || *has_device_name == 1) { return std::nullopt; } diff --git a/connections/implementation/mediums/advertisements/advertisement_util_test.cc b/connections/implementation/mediums/advertisements/advertisement_util_test.cc index e352401f..5916417d 100644 --- a/connections/implementation/mediums/advertisements/advertisement_util_test.cc +++ b/connections/implementation/mediums/advertisements/advertisement_util_test.cc @@ -24,7 +24,7 @@ namespace { TEST(AdvertisementUtilTest, ReadDeviceName) { ByteArray endpoint_info{ - "\x32\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x0b" + "\x22\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x0b" "\x54\x65\x73\x74\x20\x64\x65\x76\x69\x63\x65", 29}; std::optional parsed_device_name = ReadDeviceName(endpoint_info); diff --git a/connections/implementation/mediums/advertisements/dct_advertisement.cc b/connections/implementation/mediums/advertisements/dct_advertisement.cc index 3b691529..cd09ef90 100644 --- a/connections/implementation/mediums/advertisements/dct_advertisement.cc +++ b/connections/implementation/mediums/advertisements/dct_advertisement.cc @@ -19,10 +19,10 @@ #include #include -#include "absl/random/random.h" #include "connections/implementation/mediums/advertisements/data_element.h" #include "internal/crypto_cros/hkdf.h" #include "internal/platform/byte_array.h" +#include "internal/platform/crypto.h" #include "internal/platform/logging.h" #include "internal/platform/stream_reader.h" #include "internal/platform/stream_writer.h" @@ -37,11 +37,15 @@ constexpr int kDataTypeDeviceInformation = 0x07; constexpr int kDataTypePsm = 0x04; constexpr char kServiceIdHashSalt[] = "DCT Protocol"; constexpr char kServiceIdHashInfo[] = "Service ID Hash"; +constexpr char kEndpointIdChars[] = { + 'A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'J', 'K', 'L', + 'M', 'N', 'O', 'P', 'Q', 'R', 'S', 'T', 'U', 'V', 'W', 'X', + 'Y', 'Z', '1', '2', '3', '4', '5', '6', '7', '8', '9', '0'}; } // namespace DctAdvertisement::DctAdvertisement(const std::string& service_id, - const std::string& device_name, - uint16_t psm) { + const std::string& device_name, uint16_t psm, + uint8_t dedup) { psm_ = psm; service_id_hash_ = nearby::crypto::HkdfSha256( /*secret=*/service_id, /*salt= */ kServiceIdHashSalt, @@ -55,14 +59,13 @@ DctAdvertisement::DctAdvertisement(const std::string& service_id, device_name_ = device_name; } - absl::BitGen bitgen; - dedup_ = absl::Uniform(bitgen, 0, 1 << kDedupBitSize); + dedup_ = dedup; } std::optional DctAdvertisement::Create( - const std::string& service_id, const std::string& device_name, - uint16_t psm) { - if (service_id.empty() || device_name.empty() || psm == 0) { + const std::string& service_id, const std::string& device_name, uint16_t psm, + uint8_t dedup) { + if (service_id.empty() || device_name.empty() || psm == 0 || dedup > 0x7F) { LOG(WARNING) << "Invalid arguments for creating a DCT advertisement."; return std::nullopt; } @@ -72,7 +75,7 @@ std::optional DctAdvertisement::Create( return std::nullopt; } - return DctAdvertisement(service_id, device_name, psm); + return DctAdvertisement(service_id, device_name, psm, dedup); } std::optional DctAdvertisement::Parse( @@ -141,6 +144,26 @@ std::optional DctAdvertisement::Parse( return dct_advertisement; } +std::optional DctAdvertisement::GenerateEndpointId( + uint8_t dedup, const std::string& device_name) { + if (device_name.empty() || !IsValidUtf8String(device_name) || dedup > 0x7F) { + return std::nullopt; + } + + std::string truncted_device_name = + TrunctateDeviceName(device_name, kMaxDeviceNameSize); + + truncted_device_name.append(1, dedup); + ByteArray hash_result(4); + hash_result.CopyAt(0, Crypto::Sha256(truncted_device_name)); + std::string endpoint_id; + for (const char c : hash_result) { + endpoint_id.append(1, kEndpointIdChars[c % sizeof(kEndpointIdChars)]); + } + + return endpoint_id; +} + std::string DctAdvertisement::ToData() const { StreamWriter writer; @@ -151,8 +174,10 @@ std::string DctAdvertisement::ToData() const { writer.WriteBits(2, 4); writer.WriteBits(kDataTypePsm, 4); writer.WriteUint16(psm_); - writer.WriteBits(device_name_.size() + 1, 4); - writer.WriteBits(kDataTypeDeviceInformation, 4); + // 2 bytes DE for device information. + writer.WriteBits(1, 1); + writer.WriteBits(device_name_.size() + 1, 7); + writer.WriteUint8(kDataTypeDeviceInformation); writer.WriteBits(is_device_name_truncated_ ? 1 : 0, 1); writer.WriteBits(dedup_, kDedupBitSize); writer.WriteBytes(device_name_); diff --git a/connections/implementation/mediums/advertisements/dct_advertisement.h b/connections/implementation/mediums/advertisements/dct_advertisement.h index dba80ee0..693e2697 100644 --- a/connections/implementation/mediums/advertisements/dct_advertisement.h +++ b/connections/implementation/mediums/advertisements/dct_advertisement.h @@ -29,10 +29,13 @@ class DctAdvertisement { static std::optional Create(const std::string& service_id, const std::string& device_name, - uint16_t psm); + uint16_t psm, uint8_t dedup); static std::optional Parse( const std::string& advertisement); + static std::optional GenerateEndpointId( + uint8_t dedup, const std::string& device_name); + std::string ToData() const; uint8_t GetVersion() const { return version_; } @@ -44,7 +47,7 @@ class DctAdvertisement { private: DctAdvertisement() = default; DctAdvertisement(const std::string& service_id, - const std::string& device_name, uint16_t psm); + const std::string& device_name, uint16_t psm, uint8_t dedup); static bool IsValidUtf8String(const std::string& str); static std::string TrunctateDeviceName(const std::string& device_name, int max_length); diff --git a/connections/implementation/mediums/advertisements/dct_advertisement_test.cc b/connections/implementation/mediums/advertisements/dct_advertisement_test.cc index 1b6bdd9b..4a293c17 100644 --- a/connections/implementation/mediums/advertisements/dct_advertisement_test.cc +++ b/connections/implementation/mediums/advertisements/dct_advertisement_test.cc @@ -24,7 +24,7 @@ namespace { TEST(DctAdvertisementTest, Create) { std::optional advertisement = - DctAdvertisement::Create("service_id", "device", 0x1234); + DctAdvertisement::Create("service_id", "device", 0x1234, 0x01); ASSERT_TRUE(advertisement.has_value()); EXPECT_EQ(advertisement->GetVersion(), DctAdvertisement::kVersion); EXPECT_EQ(advertisement->GetPsm(), 0x1234); @@ -33,9 +33,14 @@ TEST(DctAdvertisementTest, Create) { } TEST(DctAdvertisementTest, CreateWithInvalidParameters) { - EXPECT_FALSE(DctAdvertisement::Create("service_id", "", 0x1234).has_value()); - EXPECT_FALSE(DctAdvertisement::Create("", "device", 0x1234).has_value()); - EXPECT_FALSE(DctAdvertisement::Create("service_id", "device", 0).has_value()); + EXPECT_FALSE( + DctAdvertisement::Create("service_id", "", 0x1234, 0x01).has_value()); + EXPECT_FALSE( + DctAdvertisement::Create("", "device", 0x1234, 0x01).has_value()); + EXPECT_FALSE( + DctAdvertisement::Create("service_id", "device", 0, 0x01).has_value()); + EXPECT_FALSE( + DctAdvertisement::Create("service_id", "device", 0, 0x81).has_value()); } TEST(DctAdvertisementTest, ParseWithInvalidParameters) { @@ -45,24 +50,25 @@ TEST(DctAdvertisementTest, ParseWithInvalidParameters) { TEST(DctAdvertisementTest, CreateWithTruncatedDeviceName) { std::optional advertisement = DctAdvertisement::Create( - "service_id", "\xC3\xA9\xC3\xB1\xC3\xB6\xF0\x9F\x98\x80", 0x1234); + "service_id", "\xC3\xA9\xC3\xB1\xC3\xB6\xF0\x9F\x98\x80", 0x1234, 0x01); EXPECT_EQ(advertisement->GetDeviceName(), "\xC3\xA9\xC3\xB1\xC3\xB6"); advertisement = DctAdvertisement::Create( - "service_id", "\xF0\x9F\x98\x80\xC3\xA9\xC3\xB1\xC3\xB6", 0x1234); + "service_id", "\xF0\x9F\x98\x80\xC3\xA9\xC3\xB1\xC3\xB6", 0x1234, 0x01); EXPECT_EQ(advertisement->GetDeviceName(), "\xF0\x9F\x98\x80\xC3\xA9"); - advertisement = DctAdvertisement::Create("service_id", "abcdefghi", 0x1234); + advertisement = + DctAdvertisement::Create("service_id", "abcdefghi", 0x1234, 0x01); EXPECT_EQ(advertisement->GetDeviceName(), "abcdefg"); advertisement = DctAdvertisement::Create( - "service_id", "\xF0\x9F\x98\x80\xF0\x9F\xAA\xB4", 0x1234); + "service_id", "\xF0\x9F\x98\x80\xF0\x9F\xAA\xB4", 0x1234, 0x01); EXPECT_EQ(advertisement->GetDeviceName(), "\xF0\x9F\x98\x80"); advertisement = DctAdvertisement::Create( - "service_id", "\xC3\xA9\xC3\xB1\xC3\xB6\xF0\x9F\x98", 0x1234); + "service_id", "\xC3\xA9\xC3\xB1\xC3\xB6\xF0\x9F\x98", 0x1234, 0x01); EXPECT_FALSE(advertisement.has_value()); } TEST(DctAdvertisementTest, CreateAndParse) { std::optional advertisement = DctAdvertisement::Create( - "service_id", "\xE4\xBD\xA0\xE5\xA5\xBD\xE7\x9A\x84", 0x1234); + "service_id", "\xE4\xBD\xA0\xE5\xA5\xBD\xE7\x9A\x84", 0x1234, 0x01); ASSERT_TRUE(advertisement.has_value()); EXPECT_EQ(advertisement->GetVersion(), DctAdvertisement::kVersion); EXPECT_EQ(advertisement->GetPsm(), 0x1234); @@ -78,5 +84,24 @@ TEST(DctAdvertisementTest, CreateAndParse) { EXPECT_EQ(parsed_advertisement->GetDeviceName(), "\xE4\xBD\xA0\xE5\xA5\xBD"); } +TEST(DctAdvertisementTest, GenerateEndpointId) { + std::optional endpoint_id = + DctAdvertisement::GenerateEndpointId(0x01, "device"); + ASSERT_TRUE(endpoint_id.has_value()); + EXPECT_EQ(endpoint_id.value(), "IWRE"); + std::optional endpoint_id_b = + DctAdvertisement::GenerateEndpointId(0x02, "device"); + ASSERT_TRUE(endpoint_id_b.has_value()); + EXPECT_EQ(endpoint_id_b.value(), "HEL4"); +} + +TEST(DctAdvertisementTest, GenerateEndpointIdWithInvalidParameters) { + EXPECT_FALSE(DctAdvertisement::GenerateEndpointId(0x01, "").has_value()); + EXPECT_FALSE( + DctAdvertisement::GenerateEndpointId(0xff, "device").has_value()); + EXPECT_FALSE( + DctAdvertisement::GenerateEndpointId(0x10, "device\xff").has_value()); +} + } // namespace } // namespace nearby::connections::advertisements::ble