diff --git a/Package.swift b/Package.swift index 9524da8a..65b8cb82 100644 --- a/Package.swift +++ b/Package.swift @@ -539,6 +539,7 @@ let package = Package( "internal/test/fake_timer_test.cc", "internal/test/fake_device_info_test.cc", "internal/test/fake_task_runner_test.cc", + "internal/weave/base_socket_test.cc", "internal/weave/control_packet_write_request_test.cc", "internal/weave/message_write_request_test.cc", "internal/weave/packet_test.cc", diff --git a/internal/weave/BUILD b/internal/weave/BUILD index 7a431e4f..541a0f06 100644 --- a/internal/weave/BUILD +++ b/internal/weave/BUILD @@ -1,22 +1,28 @@ cc_library( name = "weave", srcs = [ + "base_socket.cc", "message_write_request.cc", "packet.cc", "packet_sequence_number_generator.cc", "packetizer.cc", ], hdrs = [ + "base_socket.h", + "connection.h", "control_packet_write_request.h", "message_write_request.h", "packet.h", "packet_sequence_number_generator.h", "packetizer.h", + "socket_callback.h", ], deps = [ "//internal/platform:base", "//internal/platform:types", "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/functional:any_invocable", + "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", @@ -95,3 +101,19 @@ cc_test( "@com_google_googletest//:gtest_main", ], ) + +cc_test( + name = "base_socket_test", + srcs = [ + "base_socket_test.cc", + ], + deps = [ + ":weave", + "//internal/platform:base", + "//internal/platform:types", + "//internal/platform/implementation/g3", # build_cleaner: keep + "@com_github_protobuf_matchers//protobuf-matchers", + "@com_google_absl//absl/status", + "@com_google_googletest//:gtest_main", + ], +) diff --git a/internal/weave/base_socket.cc b/internal/weave/base_socket.cc new file mode 100644 index 00000000..8788ba81 --- /dev/null +++ b/internal/weave/base_socket.cc @@ -0,0 +1,313 @@ +// Copyright 2023 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/weave/base_socket.h" + +#include +#include +#include + +#include "absl/base/thread_annotations.h" +#include "absl/status/status.h" +#include "absl/status/statusor.h" +#include "absl/strings/str_format.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/future.h" +#include "internal/platform/logging.h" +#include "internal/platform/mutex.h" +#include "internal/platform/mutex_lock.h" +#include "internal/weave/connection.h" +#include "internal/weave/control_packet_write_request.h" +#include "internal/weave/message_write_request.h" +#include "internal/weave/packet.h" +#include "internal/weave/socket_callback.h" + +namespace nearby { +namespace weave { + +BaseSocket::BaseSocket(const Connection& connection, SocketCallback&& callback) + : socket_callback_(std::move(callback)), + connection_(const_cast(connection)) { + connection_.Initialize( + {.on_transmit_cb = + [this](absl::Status status) { + OnWriteRequestWriteComplete(status); + if (!status.ok()) { + DisconnectInternal(status); + } + }, + .on_remote_transmit_cb = + [this](std::string message) { + if (message.empty()) { + DisconnectInternal(absl::InvalidArgumentError("Empty packet!")); + return; + } + absl::StatusOr packet{ + Packet::FromBytes(ByteArray(message))}; + if (!packet.ok()) { + DisconnectInternal(packet.status()); + return; + } + bool isRemotePacketCounterExpected = + IsRemotePacketCounterExpected(packet->GetPacketCounter()); + if (packet->IsControlPacket()) { + if (!isRemotePacketCounterExpected) { + // Increment the counter by 1, ignoring the result. + remote_packet_counter_generator_.Next(); + } + OnReceiveControlPacket(std::move(*packet)); + } else { + if (isRemotePacketCounterExpected) { + OnReceiveDataPacket(std::move(*packet)); + } else { + Disconnect(); + } + } + }, + .on_disconnected_cb = [this]() { DisconnectQuietly(); }}); +} + +BaseSocket::~BaseSocket() { + NEARBY_LOGS(INFO) << "~BaseSocket"; + ShutDown(); +} + +void BaseSocket::ShutDown() { + executor_.Shutdown(); + NEARBY_LOGS(INFO) << "BaseSocket gone."; +} + +void BaseSocket::TryWriteNextControl() { + bool connected = IsConnected(); + MutexLock lock(&mutex_); + if (current_control_ == nullptr) { + if (!control_request_queue_.empty()) { + current_control_ = &control_request_queue_.front(); + } + } + if (current_control_ == nullptr && connected) { + return; + } + // We need to do this because if a control packet is being sent, it is + // one of three packets. ConnectionRequest, ConnectionConfirm, or Error. + // In any case, we should not have any messages in the queue from the previous + // connection. + current_message_ = nullptr; + message_request_queue_.clear(); + WritePacket(current_control_->NextPacket(max_packet_size_)); +} + +void BaseSocket::TryWriteNextMessage() { + bool control_queue_empty = false; + { + MutexLock lock(&mutex_); + control_queue_empty = control_request_queue_.empty(); + } + if (current_control_ != nullptr || !control_queue_empty) { + // We should only be writing one packet. + TryWriteNextControl(); + return; + } + bool connected = IsConnected(); + MutexLock lock(&mutex_); + if (current_message_ == nullptr) { + if (!message_request_queue_.empty()) { + current_message_ = &message_request_queue_.front(); + } + } + if (current_message_ == nullptr || current_message_->IsFinished() || + !connected) { + return; + } + WritePacket(current_message_->NextPacket(max_packet_size_)); +} + +void BaseSocket::WritePacket(absl::StatusOr packet) { + if (!packet.ok()) { + NEARBY_LOGS(WARNING) << "Packet status:" << packet.status(); + return; + } + CHECK(packet->SetPacketCounter(packet_counter_generator_.Next()).ok()); + NEARBY_LOGS(INFO) << "transmitting packet"; + connection_.Transmit(packet->GetBytes()); +} + +void BaseSocket::OnWriteRequestWriteComplete(absl::Status status) { + RunOnSocketThread( + "OnWriteRequestWriteComplete", + [this, status]() ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) + ABSL_LOCKS_EXCLUDED(mutex_) mutable { + { + MutexLock lock(&mutex_); + if (current_control_ != nullptr) { + current_control_ = nullptr; + control_request_queue_.pop_front(); + } else if (current_message_ != nullptr) { + NEARBY_LOGS(INFO) << "OnWriteResult current is not null"; + if (current_message_->IsFinished()) { + NEARBY_LOGS(INFO) << "OnWriteResult current finished"; + current_message_->SetWriteStatus(status); + if (!message_request_queue_.empty() && + current_message_ == &message_request_queue_.front()) { + NEARBY_LOGS(INFO) << "remove message"; + message_request_queue_.pop_front(); + current_message_ = nullptr; + } + } + } + } + TryWriteNextMessage(); + }); +} + +bool BaseSocket::IsRemotePacketCounterExpected(int counter) { + int expectedPacketCounter = remote_packet_counter_generator_.Next(); + if (counter == expectedPacketCounter) { + return true; + } + socket_callback_.on_error_cb(absl::DataLossError( + absl::StrFormat("expected remote packet counter %d for packet but got %d", + expectedPacketCounter, counter))); + return false; +} + +void BaseSocket::Disconnect() { + RunOnSocketThread("Disconnect", [this]() { + bool is_disconnecting_or_disconnected = false; + { + MutexLock lock(&mutex_); + is_disconnecting_or_disconnected = + state_ == SocketConnectionState::kDisconnecting || + state_ == SocketConnectionState::kDisconnected; + } + if (!is_disconnecting_or_disconnected) { + WriteControlPacket(Packet::CreateErrorPacket()); + { + MutexLock lock(&mutex_); + current_message_ = nullptr; + message_request_queue_.clear(); + state_ = SocketConnectionState::kDisconnecting; + } + DisconnectQuietly(); + } + }); +} + +void BaseSocket::DisconnectQuietly() { + RunOnSocketThread("ResetDisconnectQuietly", + [this]() ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) + ABSL_LOCKS_EXCLUDED(mutex_) { + bool was_connected = false; + { + MutexLock lock(&mutex_); + was_connected = + state_ == SocketConnectionState::kDisconnecting; + } + + if (was_connected) { + socket_callback_.on_disconnected_cb(); + } + packetizer_.Reset(); + packet_counter_generator_.Reset(); + remote_packet_counter_generator_.Reset(); + // Dump message and control queue. + { + MutexLock lock(&mutex_); + message_request_queue_.clear(); + control_request_queue_.clear(); + current_control_ = nullptr; + current_message_ = nullptr; + state_ = SocketConnectionState::kDisconnected; + } + NEARBY_LOGS(INFO) << "Socket now disconnected."; + }); + NEARBY_LOGS(INFO) << "scheduled reset"; +} + +void BaseSocket::OnReceiveDataPacket(Packet packet) { + absl::Status packet_status = packetizer_.AddPacket(std::move(packet)); + if (!packet_status.ok()) { + DisconnectInternal(packet_status); + return; + } + absl::StatusOr message = packetizer_.TakeMessage(); + if (!message.ok()) { + DisconnectInternal(message.status()); + return; + } + socket_callback_.on_receive_cb(message->string_data()); +} + +nearby::Future BaseSocket::Write(ByteArray message) { + MessageWriteRequest request = MessageWriteRequest(message.string_data()); + nearby::Future ret = request.GetWriteStatusFuture(); + + RunOnSocketThread( + "TryWriteMessage", + [&, request = std::move(request)]() + ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) mutable { + { + MutexLock lock(&mutex_); + message_request_queue_.push_back(std::move(request)); + } + TryWriteNextMessage(); + }); + return ret; +} + +void BaseSocket::WriteControlPacket(Packet packet) { + ControlPacketWriteRequest request = + ControlPacketWriteRequest(std::move(packet)); + RunOnSocketThread( + "TryWriteNextControl", + [&, request = std::move(request)]() + ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) mutable { + { + MutexLock lock(&mutex_); + control_request_queue_.push_back(std::move(request)); + } + TryWriteNextControl(); + }); + NEARBY_LOGS(INFO) << "Scheduled TryWriteControl"; +} + +void BaseSocket::DisconnectInternal(absl::Status status) { + socket_callback_.on_error_cb(status); + Disconnect(); +} + +bool BaseSocket::IsConnected() { + MutexLock lock(&mutex_); + return state_ == SocketConnectionState::kConnected; +} + +void BaseSocket::OnConnected(int new_max_packet_size) { + RunOnSocketThread("TryWriteOnConnected", + [this, new_max_packet_size]() + ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) { + max_packet_size_ = new_max_packet_size; + bool was_connected = IsConnected(); + if (!was_connected) { + { + MutexLock lock(&mutex_); + state_ = SocketConnectionState::kConnected; + } + socket_callback_.on_connected_cb(); + } + TryWriteNextMessage(); + }); +} + +} // namespace weave +} // namespace nearby diff --git a/internal/weave/base_socket.h b/internal/weave/base_socket.h new file mode 100644 index 00000000..693cbb6f --- /dev/null +++ b/internal/weave/base_socket.h @@ -0,0 +1,114 @@ +// Copyright 2023 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. + +#ifndef THIRD_PARTY_NEARBY_INTERNAL_WEAVE_BASE_SOCKET_H_ +#define THIRD_PARTY_NEARBY_INTERNAL_WEAVE_BASE_SOCKET_H_ + +#include +#include +#include + +#include "absl/base/thread_annotations.h" +#include "internal/platform/mutex.h" +#include "internal/platform/single_thread_executor.h" +#include "internal/weave/connection.h" +#include "internal/weave/control_packet_write_request.h" +#include "internal/weave/message_write_request.h" +#include "internal/weave/packet.h" +#include "internal/weave/packet_sequence_number_generator.h" +#include "internal/weave/packetizer.h" +#include "internal/weave/socket_callback.h" + +namespace nearby { +namespace weave { + +// The BaseSocket class covers all common sending logic and management of +// control and message packets and provides a convenient public API for sending +// messages. +class BaseSocket { + public: + BaseSocket(const Connection& connection, SocketCallback&& callback); + virtual ~BaseSocket(); + + bool IsConnected() ABSL_LOCKS_EXCLUDED(mutex_); + void Disconnect(); + nearby::Future Write(ByteArray message); + virtual void Connect() = 0; + + protected: + void OnConnected(int new_max_packet_size); + void DisconnectInternal(absl::Status status); + // DisconnectQuietly() runs the socket disconnection code without sending an + // error packet to the other side, and is needed for all disconnects, but + // mainly for the ConnectionCallback::on_disconnect_cb. + virtual void DisconnectQuietly(); + virtual void OnReceiveControlPacket(Packet packet) = 0; + void WriteControlPacket(Packet packet); + void OnReceiveDataPacket(Packet packet); + void RunOnSocketThread(std::string name, Runnable&& runnable) { + NEARBY_LOGS(INFO) << "RunOnSocketThread: " << name; + executor_.Execute(name, std::move(runnable)); + } + void ShutDown(); + + // Test only functions. + void AddControlPacket(Packet packet) ABSL_LOCKS_EXCLUDED(mutex_) { + MutexLock lock(&mutex_); + ControlPacketWriteRequest request(std::move(packet)); + control_request_queue_.push_back(std::move(request)); + } + + const Connection& GetConnection() const { return connection_; } + const SocketCallback& GetSocketCallback() const { return socket_callback_; } + + private: + enum class SocketConnectionState { + kDisconnected, + kDisconnecting, + kConnected + }; + + bool IsRemotePacketCounterExpected(int counter); + void TryWriteNextControl() ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) + ABSL_LOCKS_EXCLUDED(mutex_); + void TryWriteNextMessage() ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) + ABSL_LOCKS_EXCLUDED(mutex_); + void OnWriteRequestWriteComplete(absl::Status status) + ABSL_LOCKS_EXCLUDED(executor_); + void WritePacket(absl::StatusOr packet); + + Mutex mutex_; + // Messages and controls are in two separate queues to separate their control + // flow and to make it easier to follow the logic of sending packets. + std::deque control_request_queue_ + ABSL_GUARDED_BY(mutex_); + std::deque message_request_queue_ + ABSL_GUARDED_BY(mutex_); + ControlPacketWriteRequest* current_control_ = nullptr; + MessageWriteRequest* current_message_ = nullptr; + SocketConnectionState state_ ABSL_GUARDED_BY(mutex_) = + SocketConnectionState::kDisconnected; + int max_packet_size_; + Packetizer packetizer_; + PacketSequenceNumberGenerator packet_counter_generator_; + PacketSequenceNumberGenerator remote_packet_counter_generator_; + SocketCallback socket_callback_; + Connection& connection_; + SingleThreadExecutor executor_; +}; + +} // namespace weave +} // namespace nearby + +#endif // THIRD_PARTY_NEARBY_INTERNAL_WEAVE_BASE_SOCKET_H_ diff --git a/internal/weave/base_socket_test.cc b/internal/weave/base_socket_test.cc new file mode 100644 index 00000000..ff613432 --- /dev/null +++ b/internal/weave/base_socket_test.cc @@ -0,0 +1,478 @@ +// Copyright 2023 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/weave/base_socket.h" + +#include +#include +#include + +#include "gmock/gmock.h" +#include "protobuf-matchers/protocol-buffer-matchers.h" +#include "gtest/gtest.h" +#include "absl/status/status.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/logging.h" +#include "internal/platform/mutex.h" +#include "internal/platform/mutex_lock.h" +#include "internal/weave/connection.h" +#include "internal/weave/packet.h" +#include "internal/weave/socket_callback.h" + +namespace nearby { +namespace weave { +namespace { + +constexpr int kMaxPacketSize = 3; + +class FakeConnection : public Connection { + public: + explicit FakeConnection(int max_packet_size) + : max_packet_size_(max_packet_size) {} + void Initialize(ConnectionCallback callback) override { + callback_ = std::move(callback); + } + + int GetMaxPacketSize() override { return max_packet_size_; } + void Transmit(std::string packet) override { + packets_written_.push_back(packet); + if (instant_transmit_) { + callback_.on_transmit_cb(absl::OkStatus()); + } + } + void Close() override { open_ = false; } + bool IsOpen() { return open_; } + std::string PollWrittenPacket() { + if (!NoMorePackets()) { + auto front = packets_written_.front(); + packets_written_.erase(packets_written_.begin()); + return front; + } + NEARBY_LOGS(WARNING) << "No more packets"; + return ""; + } + bool NoMorePackets() { return packets_written_.empty(); } + void SetInstantTransmit(bool instant_transmit) { + instant_transmit_ = instant_transmit; + } + void OnTransmitProxy(absl::Status status) { + callback_.on_transmit_cb(status); + } + void OnRemoteTransmitProxy(absl::string_view message) { + callback_.on_remote_transmit_cb(std::string(message)); + } + + protected: + int max_packet_size_; + ConnectionCallback callback_; + std::vector packets_written_; + bool instant_transmit_ = true; + bool open_ = false; +}; + +class FakeSocket : public BaseSocket { + public: + explicit FakeSocket(const Connection& connection, + SocketCallback&& socketCallback) + : BaseSocket(connection, std::move(socketCallback)) {} + MOCK_METHOD(void, Connect, (), (override)); + void OnReceiveControlPacket(Packet packet) override { + control_packets_.push_back(std::move(packet)); + } + // Proxies to internal protected methods + void OnConnectedProxy(int max_packet_size) { OnConnected(max_packet_size); } + void DisconnectQuietlyProxy() { DisconnectQuietly(); } + void WriteControlPacketProxy(Packet packet) { + return WriteControlPacket(std::move(packet)); + } + void AddControlPacketToQueue(Packet packet) { + AddControlPacket(std::move(packet)); + } + + std::vector control_packets_; +}; + +Packet CreateDataPacket(int counter, bool first, bool last, ByteArray data) { + Packet packet = Packet::CreateDataPacket(first, last, data); + EXPECT_OK(packet.SetPacketCounter(counter)); + return packet; +} + +class BaseSocketTest : public ::testing::Test { + public: + BaseSocketTest() + : connection_(FakeConnection(20)), + socket_(connection_, SocketCallback{ + .on_connected_cb = + [this]() { + MutexLock lock(&mutex_); + connected_ = true; + }, + .on_disconnected_cb = + [this]() { + MutexLock lock(&mutex_); + connected_ = false; + }, + .on_receive_cb = + [this](std::string message) { + MutexLock lock(&mutex_); + messages_read_.push_back(message); + }, + .on_error_cb = + [this](absl::Status status) { + error_status_ = status; + NEARBY_LOGS(ERROR) << status; + }, + }) {} + void TransmitAndFail() { + socket_.Write(ByteArray("\x01")); + connection_.OnTransmitProxy(absl::UnavailableError("")); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + CreateDataPacket(0, true, true, ByteArray("\x01")).GetBytes()); + Packet err = Packet::CreateErrorPacket(); + EXPECT_OK(err.SetPacketCounter(1)); + EXPECT_EQ(connection_.PollWrittenPacket(), err.GetBytes()); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kUnavailable); + // clear + error_status_ = absl::OkStatus(); + } + + protected: + Mutex mutex_; + FakeConnection connection_; + FakeSocket socket_; + bool connected_ = true; + std::vector messages_read_; + absl::Status error_status_; +}; + +TEST_F(BaseSocketTest, TestConnectQueuedWrite) { + nearby::Future result = + socket_.Write(ByteArray("\x01\x02\x03")); + EXPECT_FALSE(result.IsSet()); + socket_.OnConnectedProxy(kMaxPacketSize); + // sleep for 10 ms to allow for status population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_OK(result.Get().GetResult()); + std::string packet = connection_.PollWrittenPacket(); + std::string second = connection_.PollWrittenPacket(); + Packet expected = + Packet::CreateDataPacket(true, false, ByteArray("\x01\x02")); + EXPECT_EQ(packet, expected.GetBytes()); + Packet expected2 = Packet::CreateDataPacket(false, true, ByteArray("\x03")); + EXPECT_OK(expected2.SetPacketCounter(1)); + EXPECT_EQ(second, expected2.GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestDisconnect) { + socket_.OnConnectedProxy(kMaxPacketSize); + socket_.Disconnect(); + // sleep for 10 ms to allow for executor run to complete + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + Packet::CreateErrorPacket().GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); + EXPECT_FALSE(connected_); +} + +TEST_F(BaseSocketTest, TestDisconnectTwice) { + socket_.OnConnectedProxy(kMaxPacketSize); + socket_.Disconnect(); + absl::SleepFor(absl::Milliseconds(10)); + socket_.Disconnect(); + // sleep for 10 ms to allow for executor run to complete + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + Packet::CreateErrorPacket().GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); + EXPECT_FALSE(connected_); +} + +TEST_F(BaseSocketTest, DisconnectWithoutConnectDoesNothing) { + socket_.Disconnect(); + // sleep for 10 ms to allow for executor run to complete + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_TRUE(connection_.NoMorePackets()); + // defaults to true, and the disconnect cb should not have run. + EXPECT_TRUE(connected_); +} + +TEST_F(BaseSocketTest, TestOnDisconnected) { + socket_.OnConnectedProxy(kMaxPacketSize); + // sleep for 10 ms to allow for connection status to propagate + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_TRUE(socket_.IsConnected()); + socket_.DisconnectQuietlyProxy(); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_FALSE(socket_.IsConnected()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestWriteControlPacket) { + socket_.WriteControlPacketProxy(Packet::CreateErrorPacket()); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + Packet::CreateErrorPacket().GetBytes()); + Packet second = Packet::CreateErrorPacket(); + EXPECT_OK(second.SetPacketCounter(1)); + socket_.WriteControlPacketProxy(Packet::CreateErrorPacket()); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), second.GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestWriteOnePacket) { + socket_.OnConnectedProxy(kMaxPacketSize); + nearby::Future status = socket_.Write(ByteArray("\x01\x02")); + EXPECT_OK(status.Get().GetResult()); + EXPECT_EQ( + connection_.PollWrittenPacket(), + Packet::CreateDataPacket(true, true, ByteArray("\x01\x02")).GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestWriteThreePackets) { + socket_.OnConnectedProxy(kMaxPacketSize); + nearby::Future status = + socket_.Write(ByteArray("\x01\x02\x03\x04\x05\x06")); + EXPECT_OK(status.Get().GetResult()); + EXPECT_EQ( + connection_.PollWrittenPacket(), + Packet::CreateDataPacket(true, false, ByteArray("\x01\x02")).GetBytes()); + Packet second = Packet::CreateDataPacket(false, false, ByteArray("\x03\x04")); + EXPECT_OK(second.SetPacketCounter(1)); + EXPECT_EQ(connection_.PollWrittenPacket(), second.GetBytes()); + Packet third = Packet::CreateDataPacket(false, true, ByteArray("\x05\x06")); + EXPECT_OK(third.SetPacketCounter(2)); + EXPECT_EQ(connection_.PollWrittenPacket(), third.GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestWritePacketCounterRollover) { + socket_.OnConnectedProxy(kMaxPacketSize); + for (int i = 0; i <= Packet::kMaxPacketCounter; i++) { + Packet packet = Packet::CreateDataPacket(true, true, ByteArray("\x01")); + EXPECT_OK(packet.SetPacketCounter(i)); + nearby::Future result = socket_.Write(ByteArray("\x01")); + NEARBY_LOGS(INFO) << "sent packet " << i; + EXPECT_OK(result.Get().GetResult()); + EXPECT_EQ(connection_.PollWrittenPacket(), packet.GetBytes()); + } + Packet packet = Packet::CreateDataPacket(true, true, ByteArray("\x01")); + nearby::Future result = socket_.Write(ByteArray("\x01")); + EXPECT_OK(result.Get().GetResult()); + EXPECT_EQ(connection_.PollWrittenPacket(), packet.GetBytes()); +} + +TEST_F(BaseSocketTest, TestResetByDisconnect) { + connection_.SetInstantTransmit(false); + socket_.OnConnectedProxy(kMaxPacketSize); + nearby::Future status = + socket_.Write(ByteArray("\x01\x02\x03")); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ( + connection_.PollWrittenPacket(), + Packet::CreateDataPacket(true, false, ByteArray("\x01\x02")).GetBytes()); + ASSERT_TRUE(connection_.NoMorePackets()); + + // disconnect should cause resets + socket_.Disconnect(); + EXPECT_FALSE(status.IsSet()); + EXPECT_TRUE(connection_.NoMorePackets()); + + // packet [1, 2] success + connection_.OnTransmitProxy(absl::OkStatus()); + EXPECT_FALSE(status.IsSet()); + Packet error = Packet::CreateErrorPacket(); + EXPECT_OK(error.SetPacketCounter(1)); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), error.GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); + + // error packet success + connection_.OnTransmitProxy(absl::OkStatus()); + ASSERT_TRUE(connection_.NoMorePackets()); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + + // connect again + socket_.OnConnectedProxy(kMaxPacketSize); + NEARBY_LOGS(INFO) << "Reconnected socket"; + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), ""); +} + +TEST_F(BaseSocketTest, TestResetByControlPacket) { + connection_.SetInstantTransmit(false); + socket_.OnConnectedProxy(kMaxPacketSize); + nearby::Future status = + socket_.Write(ByteArray("\x01\x02\x03")); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + CreateDataPacket(0, true, false, ByteArray("\x01\x02")).GetBytes()); + socket_.WriteControlPacketProxy(Packet::CreateErrorPacket()); + ASSERT_TRUE(connection_.NoMorePackets()); + // packet [1, 2 success] + connection_.OnTransmitProxy(absl::OkStatus()); + Packet error = Packet::CreateErrorPacket(); + ASSERT_OK(error.SetPacketCounter(1)); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), error.GetBytes()); + EXPECT_TRUE(connection_.NoMorePackets()); + // error success + connection_.OnTransmitProxy(absl::OkStatus()); + // sleep for 10 ms to allow for packet population + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestOnTransmitFailure) { + connection_.SetInstantTransmit(false); + socket_.OnConnectedProxy(kMaxPacketSize); + socket_.Write(ByteArray("\x00")); + connection_.OnTransmitProxy(absl::InternalError("EOF")); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_THAT(error_status_, + testing::status::StatusIs(absl::StatusCode::kInternal)); + EXPECT_FALSE(connected_); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitOnePacket) { + socket_.OnConnectedProxy(kMaxPacketSize); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, true, true, ByteArray("\x01\x02")).GetBytes()); + ASSERT_EQ(1, messages_read_.size()); + EXPECT_EQ(messages_read_[0], "\x01\x02"); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitTwoPackets) { + socket_.OnConnectedProxy(kMaxPacketSize); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, true, false, ByteArray("\x01\x02")).GetBytes()); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(1, false, true, ByteArray("\x03")).GetBytes()); + ASSERT_EQ(1, messages_read_.size()); + EXPECT_EQ(messages_read_[0], "\x01\x02\x03"); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitThreePackets) { + socket_.OnConnectedProxy(kMaxPacketSize); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, true, false, ByteArray("\x01\x02")).GetBytes()); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(1, false, false, ByteArray("\x03\x04")).GetBytes()); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(2, false, true, ByteArray("\x05")).GetBytes()); + ASSERT_EQ(1, messages_read_.size()); + EXPECT_EQ(messages_read_[0], "\x01\x02\x03\x04\x05"); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitPacketCounterRollover) { + for (int i = 0; i <= Packet::kMaxPacketCounter; i++) { + connection_.OnRemoteTransmitProxy( + CreateDataPacket(i, true, true, ByteArray("\x01")).GetBytes()); + ASSERT_EQ(1, messages_read_.size()); + EXPECT_EQ(messages_read_[0], "\x01"); + messages_read_.pop_back(); + } + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, true, true, ByteArray("\x01")).GetBytes()); + ASSERT_EQ(1, messages_read_.size()); + EXPECT_EQ(messages_read_[0], "\x01"); + messages_read_.pop_back(); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitIllegalDataPacket) { + socket_.OnConnectedProxy(kMaxPacketSize); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, false, false, ByteArray("\x00")).GetBytes()); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kInvalidArgument); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitDataPacketWrongPacketCounter) { + connection_.OnRemoteTransmitProxy( + CreateDataPacket(1, true, false, ByteArray("\x00")).GetBytes()); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kDataLoss); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransmitDataPacketOutOfOrderPacketCounter) { + connection_.OnRemoteTransmitProxy( + CreateDataPacket(0, true, false, ByteArray("\x00")).GetBytes()); + connection_.OnRemoteTransmitProxy( + CreateDataPacket(2, false, true, ByteArray("\x00")).GetBytes()); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kDataLoss); +} + +TEST_F(BaseSocketTest, TestOnRemoteTransitEmpty) { + connection_.OnRemoteTransmitProxy(""); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kInvalidArgument); +} + +TEST_F(BaseSocketTest, TestReconnect) { + connection_.SetInstantTransmit(false); + socket_.OnConnectedProxy(kMaxPacketSize); + NEARBY_LOGS(INFO) << "Starting TransmitAndFail1"; + TransmitAndFail(); + NEARBY_LOGS(INFO) << "TransmitAndFail1 completed"; + connection_.OnTransmitProxy(absl::UnavailableError("")); + EXPECT_EQ(error_status_.code(), absl::StatusCode::kUnavailable); + absl::SleepFor(absl::Milliseconds(20)); + NEARBY_LOGS(INFO) << "Reconnecting"; + socket_.OnConnectedProxy(kMaxPacketSize); + absl::SleepFor(absl::Milliseconds(20)); + NEARBY_LOGS(INFO) << "Starting transmit and fail 2"; + TransmitAndFail(); + NEARBY_LOGS(INFO) << "TransmitAndFail2 completed"; + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestControlGoesBeforeMessage) { + connection_.SetInstantTransmit(false); + socket_.OnConnectedProxy(kMaxPacketSize); + absl::SleepFor(absl::Milliseconds(10)); + socket_.AddControlPacketToQueue(Packet::CreateErrorPacket()); + auto future = socket_.Write(ByteArray("\x01")); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_EQ(connection_.PollWrittenPacket(), + Packet::CreateErrorPacket().GetBytes()); + EXPECT_FALSE(future.IsSet()); + EXPECT_TRUE(connection_.NoMorePackets()); + connection_.OnTransmitProxy(absl::OkStatus()); + absl::SleepFor(absl::Milliseconds(20)); + EXPECT_FALSE(future.IsSet()); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(BaseSocketTest, TestDisconnectOnBadDataPacketCounter) { + socket_.OnConnectedProxy(kMaxPacketSize); + absl::SleepFor(absl::Milliseconds(10)); + Packet bad_packet = CreateDataPacket(2, true, true, ByteArray("\x01")); + connection_.OnRemoteTransmitProxy(bad_packet.GetBytes()); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_FALSE(socket_.IsConnected()); + EXPECT_FALSE(connected_); +} + +} // namespace +} // namespace weave +} // namespace nearby diff --git a/internal/weave/connection.h b/internal/weave/connection.h new file mode 100644 index 00000000..1a05baa4 --- /dev/null +++ b/internal/weave/connection.h @@ -0,0 +1,46 @@ +// Copyright 2023 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. + +#ifndef THIRD_PARTY_NEARBY_INTERNAL_WEAVE_CONNECTION_H_ +#define THIRD_PARTY_NEARBY_INTERNAL_WEAVE_CONNECTION_H_ + +#include + +#include "absl/functional/any_invocable.h" +#include "absl/status/status.h" + +// TODO(b/269783814): Refactor into BleV2Connection class and get rid of +// this interface. +namespace nearby { +namespace weave { + +struct ConnectionCallback { + absl::AnyInvocable on_transmit_cb; + absl::AnyInvocable on_remote_transmit_cb; + absl::AnyInvocable on_disconnected_cb; +}; + +class Connection { + public: + virtual ~Connection() = default; + virtual void Initialize(ConnectionCallback callback) = 0; + virtual int GetMaxPacketSize() = 0; + virtual void Transmit(std::string packet) = 0; + virtual void Close() = 0; +}; + +} // namespace weave +} // namespace nearby + +#endif // THIRD_PARTY_NEARBY_INTERNAL_WEAVE_CONNECTION_H_ diff --git a/internal/weave/socket_callback.h b/internal/weave/socket_callback.h new file mode 100644 index 00000000..45ba024a --- /dev/null +++ b/internal/weave/socket_callback.h @@ -0,0 +1,45 @@ +// Copyright 2023 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. + +#ifndef THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKET_H_ +#define THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKET_H_ + +#include +#include + +#include "absl/status/status.h" +#include "internal/platform/logging.h" + +namespace nearby { +namespace weave { + +struct SocketCallback { + std::function on_connected_cb = []() { + NEARBY_LOGS(WARNING) << "Unimplemented!"; + }; + std::function on_disconnected_cb = []() { + NEARBY_LOGS(WARNING) << "Unimplemented!"; + }; + std::function on_receive_cb = [](std::string) { + NEARBY_LOGS(WARNING) << "Unimplemented!"; + }; + std::function on_error_cb = [](absl::Status) { + NEARBY_LOGS(WARNING) << "Unimplemented!"; + }; +}; + +} // namespace weave +} // namespace nearby + +#endif // THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKET_H_