From 9ce1286bfa0221f092d38fcacb39a7f07a289279 Mon Sep 17 00:00:00 2001 From: Edwin Wu Date: Tue, 18 Nov 2025 20:02:43 -0800 Subject: [PATCH] [BLEREFACTOR]: Implement DispatchPacket in BleSocket. PiperOrigin-RevId: 834079972 --- connections/implementation/mediums/ble/BUILD | 7 +- .../implementation/mediums/ble/ble_socket.cc | 80 +++++- .../implementation/mediums/ble/ble_socket.h | 15 ++ .../mediums/ble/ble_socket_test.cc | 239 +++++++++++++++++- 4 files changed, 334 insertions(+), 7 deletions(-) diff --git a/connections/implementation/mediums/ble/BUILD b/connections/implementation/mediums/ble/BUILD index 2e1bfe6e..a95faa6e 100644 --- a/connections/implementation/mediums/ble/BUILD +++ b/connections/implementation/mediums/ble/BUILD @@ -82,11 +82,16 @@ cc_test( deps = [ ":ble", ":ble_socket", + "//connections/implementation/flags:connections_flags", + "//internal/flags:nearby_flags", "//internal/platform:base", "//internal/platform:comm", + "//internal/platform:util", "//internal/platform/implementation:comm", - "//internal/platform/implementation/g3", + "//internal/platform/implementation/g3", # buildcleaner: keep + "//proto:connections_enums_cc_proto", "@com_github_protobuf_matchers//protobuf-matchers", + "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:string_view", "@com_google_googletest//:gtest_main", ], diff --git a/connections/implementation/mediums/ble/ble_socket.cc b/connections/implementation/mediums/ble/ble_socket.cc index 3c258a9f..1b82eec3 100644 --- a/connections/implementation/mediums/ble/ble_socket.cc +++ b/connections/implementation/mediums/ble/ble_socket.cc @@ -20,6 +20,7 @@ #include #include "absl/status/statusor.h" +#include "absl/strings/escaping.h" #include "absl/strings/str_cat.h" #include "connections/implementation/flags/nearby_connections_feature_flags.h" #include "connections/implementation/mediums/ble/ble_packet.h" @@ -196,8 +197,34 @@ Medium BleSocket::GetMediumLocked() const { } ExceptionOr BleSocket::DispatchPacket() { - // TODO(b/419654808): Implement this method. - return {Exception::kFailed}; + MutexLock lock(&mutex_); + if (!ble_input_stream_) { + return Exception::kFailed; + } + + ExceptionOr read_bytes = + ble_input_stream_->Read(BlePacket::kServiceIdHashLength); + while (read_bytes.ok()) { + ByteArray read_bytes_result = read_bytes.result(); + if (BlePacket::IsControlPacketBytes(read_bytes_result)) { + ExceptionOr handle_result = ProcessBleControlPacketLocked(); + if (!handle_result.ok()) { + return handle_result; + } + read_bytes = ble_input_stream_->Read(BlePacket::kServiceIdHashLength); + } else { + if (read_bytes_result != service_id_hash_) { + LOG(WARNING) + << "Received data packet with incorrect service ID hash. Expected: " + << absl::BytesToHexString(service_id_hash_.string_data()) + << ", Received: " + << absl::BytesToHexString(read_bytes_result.string_data()); + return Exception::kFailed; + } + break; + } + } + return read_bytes; } ExceptionOr BleSocket::ReadPayloadLength() { @@ -224,6 +251,55 @@ Exception BleSocket::WritePayloadLength(int payload_length) { return ble_output_stream_->WritePayloadLength(payload_length); } +ExceptionOr BleSocket::ProcessBleControlPacketLocked() { + // Read the first 4 bytes (packet block 1). + ExceptionOr read_bytes = ble_input_stream_->Read(4); + if (!read_bytes.ok()) { + return read_bytes; + } + if (read_bytes.result().size() != 4) { + return Exception::kFailed; + } + ByteArray packet_block_1 = read_bytes.result(); + + // Read the length from the 3rd byte of the packet block (0-indexed). + int packet_block_2_size = packet_block_1.data()[3]; + // Read the left bytes for the packet block 2). + read_bytes = ble_input_stream_->Read(packet_block_2_size); + if (!read_bytes.ok()) { + return read_bytes; + } + if (read_bytes.result().size() != packet_block_2_size) { + return Exception::kFailed; + } + ByteArray packet_block_2 = read_bytes.result(); + + // Concatenate the two packet blocks. + std::string str1(packet_block_1); + std::string str2(packet_block_2); + std::string result_str = absl::StrCat(str1, str2); + ByteArray packet_block = ByteArray(result_str); + + // Create the BlePacket from the concatenated packet block. + absl::StatusOr ble_packet_status_or = + BlePacket::CreateControlPacket(packet_block); + if (!ble_packet_status_or.ok()) { + return Exception::kFailed; + } + BlePacket ble_packet = ble_packet_status_or.value(); + ble_packet.ParseControlPacketData(packet_block.AsStringView()); + if (!ble_packet.IsValid()) { + return Exception::kFailed; + } + if (service_id_hash_ != ble_packet.GetServiceIdHash()) { + return Exception::kFailed; + } + LOG(INFO) << "Received BLE Socket Control frame: " + << BlePacket::SocketControlFrameTypeToString( + ble_packet.GetControlFrameType()); + return Exception::kSuccess; +} + Exception BleSocket::SendIntroduction() { // TODO(b/419654808): Implement this method. return {Exception::kFailed}; diff --git a/connections/implementation/mediums/ble/ble_socket.h b/connections/implementation/mediums/ble/ble_socket.h index 542a0458..38616e87 100644 --- a/connections/implementation/mediums/ble/ble_socket.h +++ b/connections/implementation/mediums/ble/ble_socket.h @@ -341,6 +341,21 @@ class BleSocket final { std::unique_ptr ble_output_stream, nearby::BleL2capSocket l2cap_socket); + /** + * Processes a control packet received from the socket. + * + * This function is called when a control packet is received from the socket. + * It is responsible for parsing the packet to identify its type and + * determining the appropriate action. + * + * @return An `ExceptionOr` containing the payload of the + * received packet on success. For certain control packets that have no + * payload, the `ByteArray` may be empty. Returns an `Exception` if a + * protocol error occurs or the read operation fails. + */ + ExceptionOr ProcessBleControlPacketLocked() + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + ::location::nearby::proto::connections::Medium GetMediumLocked() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); Exception CloseLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); diff --git a/connections/implementation/mediums/ble/ble_socket_test.cc b/connections/implementation/mediums/ble/ble_socket_test.cc index 78f7c1a0..da6ccb45 100644 --- a/connections/implementation/mediums/ble/ble_socket_test.cc +++ b/connections/implementation/mediums/ble/ble_socket_test.cc @@ -14,42 +14,79 @@ #include "connections/implementation/mediums/ble/ble_socket.h" +#include +#include #include #include #include #include #include "gtest/gtest.h" +#include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" +#include "connections/implementation/flags/nearby_connections_feature_flags.h" +#include "connections/implementation/mediums/ble/ble_packet.h" +#include "internal/flags/nearby_flags.h" #include "internal/platform/ble.h" #include "internal/platform/byte_array.h" +#include "internal/platform/byte_utils.h" #include "internal/platform/exception.h" #include "internal/platform/implementation/ble.h" #include "internal/platform/input_stream.h" #include "internal/platform/output_stream.h" +#include "proto/connections_enums.proto.h" namespace nearby { namespace connections { namespace mediums { namespace { +using ::location::nearby::proto::connections::Medium; + constexpr absl::string_view kServiceIdHash{"\x0a\x0b\x0c"}; class FakeInputStream : public InputStream { public: ExceptionOr Read(std::int64_t size) override { - return ExceptionOr(Exception::kIo); + if (exception_on_read_) { + return ExceptionOr(Exception::kIo); + } + if (buffer_.empty()) { + return ExceptionOr(Exception::kIo); + } + size_t read_size = std::min(static_cast(size), buffer_.size()); + ByteArray result{buffer_.data(), read_size}; + buffer_.erase(0, read_size); + return ExceptionOr{result}; } Exception Close() override { return {Exception::kSuccess}; } + + void Append(const ByteArray& data) { + absl::StrAppend(&buffer_, std::string(data)); + } + + void SetExceptionOnRead(bool exception_on_read) { + exception_on_read_ = exception_on_read; + } + + private: + std::string buffer_; + bool exception_on_read_ = false; }; class FakeOutputStream : public OutputStream { public: Exception Write(const ByteArray& data) override { + absl::StrAppend(&buffer_, std::string(data)); return {Exception::kSuccess}; } Exception Flush() override { return {Exception::kSuccess}; } Exception Close() override { return {Exception::kSuccess}; } + std::string GetPayload() { return buffer_; } + void Clear() { buffer_.clear(); } + + private: + std::string buffer_; }; class FakeBleSocketImpl : public api::ble::BleSocket { @@ -157,13 +194,109 @@ TEST_F(BleSocketBleMediumTest, GetRemotePeripheralReturnsSameInstance) { EXPECT_EQ(&peripheral1, &peripheral2); } +TEST_F(BleSocketBleMediumTest, DispatchPacketWithCorrectServiceIdHash) { + fake_input_stream_.Append(ByteArray(std::string(kServiceIdHash))); + auto result = socket_->DispatchPacket(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), ByteArray(std::string(kServiceIdHash))); +} + +TEST_F(BleSocketBleMediumTest, DispatchPacketWithWrongServiceIdHash) { + fake_input_stream_.Append(ByteArray("\x01\x02\x03")); + auto result = socket_->DispatchPacket(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kFailed); +} + +TEST_F(BleSocketBleMediumTest, DispatchPacketWithControlPacket) { + auto control_packet_status = BlePacket::CreateControlIntroductionPacket( + ByteArray(std::string(kServiceIdHash))); + ASSERT_TRUE(control_packet_status.ok()); + BlePacket control_packet = control_packet_status.value(); + fake_input_stream_.Append(ByteArray(control_packet)); + fake_input_stream_.Append(ByteArray(std::string(kServiceIdHash))); + + auto result = socket_->DispatchPacket(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), ByteArray(std::string(kServiceIdHash))); +} + +TEST_F(BleSocketBleMediumTest, DispatchPacketWithReadError) { + fake_input_stream_.SetExceptionOnRead(true); + auto result = socket_->DispatchPacket(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleSocketBleMediumTest, ReadPayloadLengthSuccess) { + constexpr int kPayloadLength = 12345; + fake_input_stream_.Append(ByteArray(byte_utils::IntToBytes(kPayloadLength))); + auto result = socket_->ReadPayloadLength(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), kPayloadLength); +} + +TEST_F(BleSocketBleMediumTest, ReadPayloadLengthReadError) { + fake_input_stream_.SetExceptionOnRead(true); + auto result = socket_->ReadPayloadLength(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleSocketBleMediumTest, ReadPayloadLengthAfterClose) { + socket_->Close(); + auto result = socket_->ReadPayloadLength(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleSocketBleMediumTest, WritePayloadLengthSuccess) { + constexpr int kPayloadLength = 12345; + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleSocketBleMediumTest, WritePayloadLengthAfterClose) { + constexpr int kPayloadLength = 12345; + socket_->Close(); + EXPECT_FALSE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleSocketBleMediumTest, WritePayloadLengthFailsIfCalledTwice) { + NearbyFlags::GetInstance().OverrideBoolFlagValue( + config_package_nearby::nearby_connections_feature::kRefactorBleL2cap, + true); + constexpr int kPayloadLength = 12345; + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); + EXPECT_FALSE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleSocketBleMediumTest, WritePayloadLengthSucceedsIfCalledAfterWrite) { + NearbyFlags::GetInstance().OverrideBoolFlagValue( + config_package_nearby::nearby_connections_feature::kRefactorBleL2cap, + true); + constexpr int kPayloadLength = 12345; + ByteArray payload("payload"); + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); + EXPECT_TRUE(socket_->GetOutputStream().Write(payload).Ok()); + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleSocketBleMediumTest, IsValid) { + EXPECT_TRUE(socket_->IsValid()); + EXPECT_TRUE(socket_->Close().Ok()); + EXPECT_TRUE(socket_->IsValid()); +} + +TEST_F(BleSocketBleMediumTest, GetMedium) { + EXPECT_EQ(socket_->GetMedium(), Medium::BLE); +} + class BleL2capSocketBleMediumTest : public ::testing::Test { protected: void SetUp() override { nearby::BlePeripheral peripheral; - auto fake_l2cap_socket_impl = - std::make_unique( - fake_input_stream_, fake_output_stream_); + auto fake_l2cap_socket_impl = std::make_unique( + fake_input_stream_, fake_output_stream_); fake_l2cap_socket_impl_ = fake_l2cap_socket_impl.get(); @@ -215,6 +348,104 @@ TEST_F(BleL2capSocketBleMediumTest, GetRemotePeripheralReturnsSameInstance) { EXPECT_EQ(&peripheral1, &peripheral2); } +TEST_F(BleL2capSocketBleMediumTest, DispatchPacketWithCorrectServiceIdHash) { + fake_input_stream_.Append(ByteArray(std::string(kServiceIdHash))); + auto result = socket_->DispatchPacket(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), ByteArray(std::string(kServiceIdHash))); +} + +TEST_F(BleL2capSocketBleMediumTest, DispatchPacketWithWrongServiceIdHash) { + fake_input_stream_.Append(ByteArray("\x01\x02\x03")); + auto result = socket_->DispatchPacket(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kFailed); +} + +TEST_F(BleL2capSocketBleMediumTest, DispatchPacketWithControlPacket) { + auto control_packet_status = BlePacket::CreateControlIntroductionPacket( + ByteArray(std::string(kServiceIdHash))); + ASSERT_TRUE(control_packet_status.ok()); + BlePacket control_packet = control_packet_status.value(); + fake_input_stream_.Append(ByteArray(control_packet)); + fake_input_stream_.Append(ByteArray(std::string(kServiceIdHash))); + + auto result = socket_->DispatchPacket(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), ByteArray(std::string(kServiceIdHash))); +} + +TEST_F(BleL2capSocketBleMediumTest, DispatchPacketWithReadError) { + fake_input_stream_.SetExceptionOnRead(true); + auto result = socket_->DispatchPacket(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleL2capSocketBleMediumTest, ReadPayloadLengthSuccess) { + constexpr int kPayloadLength = 12345; + fake_input_stream_.Append(ByteArray(byte_utils::IntToBytes(kPayloadLength))); + auto result = socket_->ReadPayloadLength(); + EXPECT_TRUE(result.ok()); + EXPECT_EQ(result.result(), kPayloadLength); +} + +TEST_F(BleL2capSocketBleMediumTest, ReadPayloadLengthReadError) { + fake_input_stream_.SetExceptionOnRead(true); + auto result = socket_->ReadPayloadLength(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleL2capSocketBleMediumTest, ReadPayloadLengthAfterClose) { + socket_->Close(); + auto result = socket_->ReadPayloadLength(); + EXPECT_FALSE(result.ok()); + EXPECT_EQ(result.exception(), Exception::kIo); +} + +TEST_F(BleL2capSocketBleMediumTest, WritePayloadLengthSuccess) { + constexpr int kPayloadLength = 12345; + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleL2capSocketBleMediumTest, WritePayloadLengthAfterClose) { + constexpr int kPayloadLength = 12345; + socket_->Close(); + EXPECT_FALSE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleL2capSocketBleMediumTest, WritePayloadLengthFailsIfCalledTwice) { + NearbyFlags::GetInstance().OverrideBoolFlagValue( + config_package_nearby::nearby_connections_feature::kRefactorBleL2cap, + true); + constexpr int kPayloadLength = 12345; + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); + EXPECT_FALSE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleL2capSocketBleMediumTest, + WritePayloadLengthSucceedsIfCalledAfterWrite) { + NearbyFlags::GetInstance().OverrideBoolFlagValue( + config_package_nearby::nearby_connections_feature::kRefactorBleL2cap, + true); + constexpr int kPayloadLength = 12345; + ByteArray payload("payload"); + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); + EXPECT_TRUE(socket_->GetOutputStream().Write(payload).Ok()); + EXPECT_TRUE(socket_->WritePayloadLength(kPayloadLength).Ok()); +} + +TEST_F(BleL2capSocketBleMediumTest, IsValid) { + EXPECT_TRUE(socket_->IsValid()); + EXPECT_TRUE(socket_->Close().Ok()); + EXPECT_TRUE(socket_->IsValid()); +} + +TEST_F(BleL2capSocketBleMediumTest, GetMedium) { + EXPECT_EQ(socket_->GetMedium(), Medium::BLE_L2CAP); +} + } // namespace } // namespace mediums } // namespace connections