refactor on DCT advertisement

PiperOrigin-RevId: 731443070
This commit is contained in:
Guogang Li
2025-02-26 13:52:29 -08:00
committed by Copybara-Service
parent 5c3cc4eb71
commit 6ff9da2d90
6 changed files with 89 additions and 27 deletions
@@ -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",
],
@@ -28,12 +28,12 @@ std::optional<std::string> ReadDeviceName(const ByteArray& endpoint_info) {
// LINT.IfChange
StreamReader reader(endpoint_info);
std::optional<uint8_t> version = reader.ReadBits(3);
if (!version.has_value() || *version != 1) {
if (!version.has_value() || *version > 1) {
return std::nullopt;
}
std::optional<uint8_t> 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;
}
@@ -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<std::string> parsed_device_name = ReadDeviceName(endpoint_info);
@@ -19,10 +19,10 @@
#include <optional>
#include <string>
#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> 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> DctAdvertisement::Create(
return std::nullopt;
}
return DctAdvertisement(service_id, device_name, psm);
return DctAdvertisement(service_id, device_name, psm, dedup);
}
std::optional<DctAdvertisement> DctAdvertisement::Parse(
@@ -141,6 +144,26 @@ std::optional<DctAdvertisement> DctAdvertisement::Parse(
return dct_advertisement;
}
std::optional<std::string> 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_);
@@ -29,10 +29,13 @@ class DctAdvertisement {
static std::optional<DctAdvertisement> Create(const std::string& service_id,
const std::string& device_name,
uint16_t psm);
uint16_t psm, uint8_t dedup);
static std::optional<DctAdvertisement> Parse(
const std::string& advertisement);
static std::optional<std::string> 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);
@@ -24,7 +24,7 @@ namespace {
TEST(DctAdvertisementTest, Create) {
std::optional<DctAdvertisement> 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<DctAdvertisement> 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<DctAdvertisement> 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<std::string> endpoint_id =
DctAdvertisement::GenerateEndpointId(0x01, "device");
ASSERT_TRUE(endpoint_id.has_value());
EXPECT_EQ(endpoint_id.value(), "IWRE");
std::optional<std::string> 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