diff --git a/internal/platform/implementation/linux/BUILD b/internal/platform/implementation/linux/BUILD index 1b55b1c4..ad1b127e 100644 --- a/internal/platform/implementation/linux/BUILD +++ b/internal/platform/implementation/linux/BUILD @@ -196,6 +196,7 @@ cc_library( "ble_v2_server_socket.cc", "ble_v2_socket.cc", "ble_l2cap_server_socket.cc", + "ble_l2cap_socket.cc", "bluetooth_adapter.cc", "bluetooth_bluez_profile.cc", "bluetooth_classic_device.cc", diff --git a/internal/platform/implementation/linux/ble_l2cap_connection_test.cc b/internal/platform/implementation/linux/ble_l2cap_connection_test.cc new file mode 100644 index 00000000..537c6fba --- /dev/null +++ b/internal/platform/implementation/linux/ble_l2cap_connection_test.cc @@ -0,0 +1,98 @@ +// Copyright 2025 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "internal/platform/implementation/linux/ble_l2cap_connection.h" + +#include +#include + +#include +#include + +#include "internal/platform/byte_array.h" +#include "gtest/gtest.h" + +namespace nearby { +namespace linux { +namespace { + +std::string WrapFrame(absl::string_view payload) { + return std::string{ + static_cast((payload.size() >> 24) & 0xFF), + static_cast((payload.size() >> 16) & 0xFF), + static_cast((payload.size() >> 8) & 0xFF), + static_cast(payload.size() & 0xFF)} + + std::string(payload); +} + +class BleL2capConnectionTest : public ::testing::Test { + protected: + void SetUp() override { + ASSERT_EQ(socketpair(AF_UNIX, SOCK_STREAM, 0, fds_), 0); + peer_fd_ = fds_[1]; + streams_ = CreateBleL2capConnectionStreams(fds_[0]); + } + + void TearDown() override { + if (streams_.input) streams_.input->Close(); + if (streams_.output) streams_.output->Close(); + if (peer_fd_ >= 0) close(peer_fd_); + } + + int fds_[2]{-1, -1}; + int peer_fd_{-1}; + BleL2capConnectionStreams streams_; +}; + +TEST_F(BleL2capConnectionTest, ReadReceivesExactBytesFromPeer) { + std::string message = "hello into l2cap socket"; + std::string framed_message = WrapFrame(message); + + ASSERT_EQ(write(peer_fd_, framed_message.data(), framed_message.size()), + static_cast(framed_message.size())); + + auto out = streams_.input->Read(message.size()).GetResult(); + EXPECT_EQ(out, ByteArray(message)); + EXPECT_EQ(out.AsStringView(), message); +} + +TEST_F(BleL2capConnectionTest, WriteSendsExactBytesToPeer) { + std::string message = "hello from l2cap socket"; + std::string framed_message = WrapFrame(message); + + EXPECT_EQ(streams_.output->Write(message).value, Exception::kSuccess); + + std::vector received(framed_message.size()); + ssize_t n = read(peer_fd_, received.data(), received.size()); + + ASSERT_EQ(n, static_cast(framed_message.size())); + EXPECT_EQ(std::string(received.begin(), received.end()), framed_message); +} + +TEST_F(BleL2capConnectionTest, SkipSkipsBytesFromPeer) { + std::string message = "hello there"; + std::string framed_message = WrapFrame(message); + + ASSERT_EQ(write(peer_fd_, framed_message.data(), framed_message.size()), + static_cast(framed_message.size())); + + EXPECT_EQ(streams_.input->Skip(6).result(), 6); + + auto out = streams_.input->Read(5).GetResult(); + EXPECT_EQ(out, ByteArray("there")); +} + +} // namespace +} // namespace linux +} // namespace nearby diff --git a/internal/platform/implementation/linux/ble_l2cap_socket.cc b/internal/platform/implementation/linux/ble_l2cap_socket.cc index ae39a4c4..7bea3192 100644 --- a/internal/platform/implementation/linux/ble_l2cap_socket.cc +++ b/internal/platform/implementation/linux/ble_l2cap_socket.cc @@ -14,28 +14,104 @@ #include "internal/platform/implementation/linux/ble_l2cap_socket.h" -#include -#include -#include - #include #include -#include -#include +#include #include -#include +#include "absl/strings/escaping.h" #include "absl/strings/string_view.h" -#include "absl/time/time.h" -#include "internal/platform/implementation/crypto.h" #include "internal/platform/logging.h" -#include "proto/mediums/ble_frames.pb.h" namespace nearby { namespace linux { -// Migrated from bespoke l2cap socket stream semantics to linux platform stream -// semantics -// +namespace { + +constexpr int kFrameHeaderSize = 4; +constexpr int kMaxFrameSize = 1024 * 1024; + +std::string EncodeFrameLength(uint32_t length) { + return std::string{ + static_cast((length >> 24) & 0xFF), + static_cast((length >> 16) & 0xFF), + static_cast((length >> 8) & 0xFF), + static_cast(length & 0xFF), + }; +} + +uint32_t DecodeFrameLength(absl::string_view header) { + return (static_cast(static_cast(header[0])) << 24) | + (static_cast(static_cast(header[1])) << 16) | + (static_cast(static_cast(header[2])) << 8) | + static_cast(static_cast(header[3])); +} + +} // namespace + +Exception BleL2capOutputStream::Write(absl::string_view data) { + if (data.size() > std::numeric_limits::max()) { + return {Exception::kIo}; + } + + std::string frame = EncodeFrameLength(static_cast(data.size())); + frame.append(data.data(), data.size()); + return stream_.Write(frame); +} + +Exception BleL2capOutputStream::Flush() { + return stream_.Flush(); +} + +ExceptionOr BleL2capInputStream::Read(std::int64_t size) { + if (size <= 0) { + return ExceptionOr(ByteArray(std::string())); + } + LOG(INFO) << __func__ << ": trying to read " << size << " payload bytes"; + + while (pending_.empty()) { + ExceptionOr packet = stream_.Read(kMaxFrameSize); + if (!packet.ok()) { + return {Exception::kIo}; + } + + std::string packet_data = packet.result().string_data(); + + if (packet_data.empty()) { + return {Exception::kIo}; + } + + wire_buffer_.append(packet_data); + + while (wire_buffer_.size() >= kFrameHeaderSize) { + uint32_t frame_length = DecodeFrameLength(absl::string_view( + wire_buffer_.data(), kFrameHeaderSize)); + if (frame_length > kMaxFrameSize) { + LOG(ERROR) << __func__ << ": invalid L2CAP frame length " + << frame_length; + return {Exception::kIo}; + } + + size_t full_frame_size = kFrameHeaderSize + frame_length; + if (wire_buffer_.size() < full_frame_size) { + break; + } + + pending_.append(wire_buffer_.data() + kFrameHeaderSize, frame_length); + wire_buffer_.erase(0, full_frame_size); + + if (!pending_.empty()) { + break; + } + + continue; + } + } + + size_t bytes_to_return = std::min(static_cast(size), pending_.size()); + std::string out = pending_.substr(0, bytes_to_return); + pending_.erase(0, bytes_to_return); + return ExceptionOr(ByteArray(std::move(out))); +} } // namespace linux } // namespace nearby diff --git a/internal/platform/implementation/linux/ble_l2cap_socket.h b/internal/platform/implementation/linux/ble_l2cap_socket.h index e05305ba..80964624 100644 --- a/internal/platform/implementation/linux/ble_l2cap_socket.h +++ b/internal/platform/implementation/linux/ble_l2cap_socket.h @@ -17,7 +17,7 @@ #include "dbus.h" -#include +#include #include #include #include @@ -28,21 +28,63 @@ #include "internal/platform/byte_array.h" #include "internal/platform/exception.h" #include "internal/platform/implementation/ble.h" - #include "internal/platform/implementation/linux/stream.h" + namespace nearby { namespace linux { class BleL2capSocket; +class BleL2capInputStream : public nearby::InputStream { + public: + BleL2capInputStream(sdbus::UnixFd fd) : stream_(fd) {} + + ExceptionOr Read(std::int64_t size) override; + Exception Close() override { + if (closed_) { + return {Exception::kSuccess}; + } + closed_ = true; + return stream_.Close(); + } + + private: + linux::InputStream stream_; + std::string wire_buffer_; + std::string pending_; + bool closed_ = false; +}; + +class BleL2capOutputStream : public nearby::OutputStream { + public: + BleL2capOutputStream(sdbus::UnixFd fd) : stream_(fd) {} + + Exception Write(absl::string_view data) override; + Exception Flush() override; + Exception Close() override { + if (closed_) { + return {Exception::kSuccess}; + } + closed_ = true; + return stream_.Close(); + } + + private: + linux::OutputStream stream_; + bool closed_ = false; +}; + class BleL2capSocket final : public api::ble::BleL2capSocket { public: BleL2capSocket(int fd, api::ble::BlePeripheral::UniqueId peripheral_id, std::string service_id = "") - : fd_(sdbus::UnixFd(fd)), output_stream_(fd_), input_stream_(fd_) {}; + : fd_(sdbus::UnixFd(fd)), + output_stream_(fd_), + input_stream_(fd_), + peripheral_id_(peripheral_id) {}; - InputStream& GetInputStream() override { return input_stream_; } - OutputStream& GetOutputStream() override { return output_stream_; } + nearby::InputStream& GetInputStream() override { return input_stream_; } + nearby::OutputStream& GetOutputStream() override { return output_stream_; } Exception Close() override { input_stream_.Close(); output_stream_.Close(); @@ -55,8 +97,8 @@ class BleL2capSocket final : public api::ble::BleL2capSocket { private: sdbus::UnixFd fd_; - OutputStream output_stream_; - InputStream input_stream_; + BleL2capOutputStream output_stream_; + BleL2capInputStream input_stream_; api::ble::BlePeripheral::UniqueId peripheral_id_; }; diff --git a/internal/platform/implementation/linux/ble_l2cap_socket_test.cc b/internal/platform/implementation/linux/ble_l2cap_socket_test.cc index fdc64ba7..8f81f8bb 100644 --- a/internal/platform/implementation/linux/ble_l2cap_socket_test.cc +++ b/internal/platform/implementation/linux/ble_l2cap_socket_test.cc @@ -31,6 +31,16 @@ namespace nearby { namespace linux { namespace { + +std::string WrapFrame(absl::string_view payload) { + return std::string{ + static_cast((payload.size() >> 24) & 0xFF), + static_cast((payload.size() >> 16) & 0xFF), + static_cast((payload.size() >> 8) & 0xFF), + static_cast(payload.size() & 0xFF)} + + std::string(payload); +} + class BleL2capSocketTest : public ::testing::Test { protected: void SetUp() override { @@ -57,31 +67,32 @@ void SetUp() override { std::unique_ptr socket_; }; TEST_F(BleL2capSocketTest, ReturnsInputAndOutputStreams) { - InputStream& input = socket_->GetInputStream(); - OutputStream& output = socket_->GetOutputStream(); + nearby::InputStream& input = socket_->GetInputStream(); + nearby::OutputStream& output = socket_->GetOutputStream(); EXPECT_NE(&input, nullptr); EXPECT_NE(&output, nullptr); } TEST_F(BleL2capSocketTest, ReturnsSameStreamInstancesAcrossCalls) { - InputStream& input1 = socket_->GetInputStream(); - InputStream& input2 = socket_->GetInputStream(); + nearby::InputStream& input1 = socket_->GetInputStream(); + nearby::InputStream& input2 = socket_->GetInputStream(); - OutputStream& output1 = socket_->GetOutputStream(); - OutputStream& output2 = socket_->GetOutputStream(); + nearby::OutputStream& output1 = socket_->GetOutputStream(); + nearby::OutputStream& output2 = socket_->GetOutputStream(); EXPECT_EQ(&input1, &input2); EXPECT_EQ(&output1, &output2); } TEST_F(BleL2capSocketTest, ReadReceivesExactBytesFromPeer) { std::string message = "hello into l2cap socket"; + std::string framed_message = WrapFrame(message); ASSERT_EQ( - write(peer_fd_, message.data(), message.size()), - static_cast(message.size()) + write(peer_fd_, framed_message.data(), framed_message.size()), + static_cast(framed_message.size()) ); - InputStream& input = socket_->GetInputStream(); + nearby::InputStream& input = socket_->GetInputStream(); std::vector buffer(message.size()); auto out = input.Read(buffer.size()).GetResult(); @@ -98,31 +109,33 @@ TEST_F(BleL2capSocketTest, ReadReceivesExactBytesFromPeer) { } TEST_F(BleL2capSocketTest, WriteSendsExactBytesToPeer) { std::string message = "hello from l2cap socket"; + std::string framed_message = WrapFrame(message); - OutputStream& output = socket_->GetOutputStream(); + nearby::OutputStream& output = socket_->GetOutputStream(); EXPECT_EQ(output.Write(message).value, Exception::kSuccess); - std::vector received(message.size()); + std::vector received(framed_message.size()); ssize_t n = read(peer_fd_, received.data(), received.size()); - ASSERT_EQ(n, static_cast(message.size())); + ASSERT_EQ(n, static_cast(framed_message.size())); EXPECT_EQ( std::string(received.begin(), received.end()), - message + framed_message ); } TEST_F(BleL2capSocketTest, SkipSkipsBytesFromPeer) { std::string message = "hello there"; + std::string framed_message = WrapFrame(message); ASSERT_EQ( - write(peer_fd_, message.data(), message.size()), - static_cast(message.size()) + write(peer_fd_, framed_message.data(), framed_message.size()), + static_cast(framed_message.size()) ); - InputStream& input = socket_->GetInputStream(); + nearby::InputStream& input = socket_->GetInputStream(); EXPECT_EQ(input.Skip(6).result(), 6); @@ -132,6 +145,26 @@ TEST_F(BleL2capSocketTest, SkipSkipsBytesFromPeer) { TEST_F(BleL2capSocketTest, CloseReturnsSuccess) { EXPECT_EQ(socket_->Close().value, Exception::kSuccess); } + +TEST(BleL2capSocketSeqpacketTest, ReadReceivesHeaderAndPayloadFromOnePacket) { + int fds[2]{-1, -1}; + ASSERT_EQ(socketpair(AF_UNIX, SOCK_SEQPACKET, 0, fds), 0); + + auto socket = std::make_unique(fds[0], 1); + int peer_fd = fds[1]; + + std::string message = "\x03"; + std::string framed_message = WrapFrame(message); + + ASSERT_EQ(write(peer_fd, framed_message.data(), framed_message.size()), + static_cast(framed_message.size())); + + auto out = socket->GetInputStream().Read(1).GetResult(); + EXPECT_EQ(out, ByteArray(message)); + + socket->Close(); + close(peer_fd); +} } // namespace } // namespace linux } // namespace nearby diff --git a/internal/platform/implementation/linux/stream.cc b/internal/platform/implementation/linux/stream.cc index bf25e427..48997f3f 100644 --- a/internal/platform/implementation/linux/stream.cc +++ b/internal/platform/implementation/linux/stream.cc @@ -20,6 +20,7 @@ #include #include +#include "absl/strings/escaping.h" #include "internal/platform/byte_array.h" #include "internal/platform/exception.h" #include "internal/platform/implementation/linux/stream.h" @@ -63,8 +64,7 @@ ExceptionOr InputStream::Read(std::int64_t size) { } if (pfd.revents & (POLLIN | POLLHUP)) { - ssize_t bytes_read = - recv(fd_->get(), buffer.data(), buffer.size(), 0); + ssize_t bytes_read = recv(fd_->get(), buffer.data(), buffer.size(), 0); if (bytes_read > 0) { buffer.resize(static_cast(bytes_read)); @@ -117,8 +117,7 @@ Exception OutputStream::Write(absl::string_view data) { } while (poll_result < 0 && errno == EINTR); if (poll_result < 0) { - LOG(ERROR) << __func__ - << ": poll failed: " << std::strerror(errno); + LOG(ERROR) << __func__ << ": poll failed: " << std::strerror(errno); return {Exception::kIo}; } @@ -132,11 +131,7 @@ Exception OutputStream::Write(absl::string_view data) { continue; } - ssize_t n = send( - fd, - data.data() + sent, - data.size() - sent, - MSG_NOSIGNAL); + ssize_t n = send(fd, data.data() + sent, data.size() - sent, MSG_NOSIGNAL); if (n > 0) { sent += static_cast(n); @@ -158,21 +153,22 @@ Exception OutputStream::Write(absl::string_view data) { continue; } - LOG(ERROR) << __func__ - << ": error writing to fd: " << std::strerror(errno); + LOG(ERROR) << __func__ << ": error writing to fd: " << std::strerror(errno); return {Exception::kIo}; } return {Exception::kSuccess}; } -Exception OutputStream::Flush() { return Exception{Exception::kSuccess}; } +Exception OutputStream::Flush() { + return Exception{Exception::kSuccess}; +} Exception OutputStream::Close() { if (!fd_->isValid()) return Exception{Exception::kIo}; auto ret = close(fd_->get()) < 0 ? Exception{Exception::kIo} - : Exception{Exception::kSuccess}; + : Exception{Exception::kSuccess}; fd_.reset(); return ret; }