From f95a875cded3707659f71036afe78c378ad031ff Mon Sep 17 00:00:00 2001 From: Anay Wadhera Date: Thu, 13 Jul 2023 08:43:31 -0700 Subject: [PATCH] Weave port [7/n]: Introduce Weave client socket implementation PiperOrigin-RevId: 547812816 --- Package.swift | 2 + internal/platform/BUILD | 4 +- internal/weave/BUILD | 3 + internal/weave/base_socket_test.cc | 2 +- internal/weave/connection.h | 2 +- internal/weave/sockets/BUILD | 34 ++ internal/weave/sockets/client_socket.cc | 125 ++++++ internal/weave/sockets/client_socket.h | 60 +++ internal/weave/sockets/client_socket_test.cc | 365 ++++++++++++++++++ .../weave/sockets/initial_data_provider.h | 49 +++ 10 files changed, 642 insertions(+), 4 deletions(-) create mode 100644 internal/weave/sockets/BUILD create mode 100644 internal/weave/sockets/client_socket.cc create mode 100644 internal/weave/sockets/client_socket.h create mode 100644 internal/weave/sockets/client_socket_test.cc create mode 100644 internal/weave/sockets/initial_data_provider.h diff --git a/Package.swift b/Package.swift index a887a0b9..cf413db8 100644 --- a/Package.swift +++ b/Package.swift @@ -398,6 +398,7 @@ let package = Package( "internal/crypto/BUILD.gn", "internal/interop/BUILD", "internal/weave/BUILD", + "internal/weave/sockets/BUILD", "internal/platform/flags/BUILD", "internal/platform/implementation/shared/BUILD", "internal/platform/implementation/apple/Mediums/BUILD", @@ -550,6 +551,7 @@ let package = Package( "internal/weave/packet_test.cc", "internal/weave/packet_sequence_number_generator_test.cc", "internal/weave/packetizer_test.cc", + "internal/weave/sockets/client_socket_test.cc", // simulation "connections/implementation/offline_simulation_user.cc", "connections/implementation/simulation_user.cc", diff --git a/internal/platform/BUILD b/internal/platform/BUILD index 09c6ada8..c6daaee4 100644 --- a/internal/platform/BUILD +++ b/internal/platform/BUILD @@ -51,7 +51,7 @@ cc_library( "//internal/platform:__subpackages__", "//internal/platform/implementation:__subpackages__", "//internal/preferences:__subpackages__", - "//internal/weave:__pkg__", + "//internal/weave:__subpackages__", "//location/nearby/cpp:__subpackages__", "//presence:__subpackages__", "//third_party/nearby/sharing:__subpackages__", @@ -363,7 +363,7 @@ cc_library( "//internal/platform/implementation/windows:__subpackages__", "//internal/preferences:__subpackages__", "//internal/test:__subpackages__", - "//internal/weave:__pkg__", + "//internal/weave:__subpackages__", "//location/nearby/cpp:__subpackages__", "//location/nearby/testing/nearby_native:__subpackages__", "//presence:__subpackages__", diff --git a/internal/weave/BUILD b/internal/weave/BUILD index 541a0f06..19b814b5 100644 --- a/internal/weave/BUILD +++ b/internal/weave/BUILD @@ -17,6 +17,9 @@ cc_library( "packetizer.h", "socket_callback.h", ], + visibility = [ + ":__subpackages__", + ], deps = [ "//internal/platform:base", "//internal/platform:types", diff --git a/internal/weave/base_socket_test.cc b/internal/weave/base_socket_test.cc index ff613432..f08ef50d 100644 --- a/internal/weave/base_socket_test.cc +++ b/internal/weave/base_socket_test.cc @@ -44,7 +44,7 @@ class FakeConnection : public Connection { callback_ = std::move(callback); } - int GetMaxPacketSize() override { return max_packet_size_; } + int GetMaxPacketSize() const override { return max_packet_size_; } void Transmit(std::string packet) override { packets_written_.push_back(packet); if (instant_transmit_) { diff --git a/internal/weave/connection.h b/internal/weave/connection.h index 1a05baa4..7269611b 100644 --- a/internal/weave/connection.h +++ b/internal/weave/connection.h @@ -35,7 +35,7 @@ class Connection { public: virtual ~Connection() = default; virtual void Initialize(ConnectionCallback callback) = 0; - virtual int GetMaxPacketSize() = 0; + virtual int GetMaxPacketSize() const = 0; virtual void Transmit(std::string packet) = 0; virtual void Close() = 0; }; diff --git a/internal/weave/sockets/BUILD b/internal/weave/sockets/BUILD new file mode 100644 index 00000000..9b6feb03 --- /dev/null +++ b/internal/weave/sockets/BUILD @@ -0,0 +1,34 @@ +cc_library( + name = "sockets", + srcs = [ + "client_socket.cc", + ], + hdrs = [ + "client_socket.h", + "initial_data_provider.h", + ], + visibility = [ + "//connections:__subpackages__", + ], + deps = [ + "//internal/weave", + "@com_google_absl//absl/random", + "@com_google_absl//absl/status", + ], +) + +cc_test( + name = "sockets_test", + srcs = [ + "client_socket_test.cc", + ], + deps = [ + ":sockets", + "//internal/platform:base", + "//internal/platform:types", + "//internal/platform/implementation/g3", # build_cleaner: keep + "@com_github_protobuf_matchers//protobuf-matchers", + "@com_google_absl//absl/time", + "@com_google_googletest//:gtest_main", + ], +) diff --git a/internal/weave/sockets/client_socket.cc b/internal/weave/sockets/client_socket.cc new file mode 100644 index 00000000..cd4ef1e5 --- /dev/null +++ b/internal/weave/sockets/client_socket.cc @@ -0,0 +1,125 @@ +// 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/sockets/client_socket.h" + +#include +#include +#include + +#include "absl/status/status.h" +#include "internal/weave/base_socket.h" +#include "internal/weave/socket_callback.h" +#include "internal/weave/sockets/initial_data_provider.h" + +namespace nearby { +namespace weave { +namespace { +constexpr int kConnectionConfirmPacketMinLength = 4; +} + +ClientSocket::ClientSocket(const Connection& connection, + SocketCallback&& socket_callback) + : BaseSocket(connection, std::move(socket_callback)) { + initial_data_provider_ = std::make_unique(); +} + +ClientSocket::ClientSocket( + const Connection& connection, SocketCallback&& socket_callback, + std::unique_ptr initial_data_provider) + : BaseSocket(connection, std::move(socket_callback)), + initial_data_provider_(std::move(initial_data_provider)) {} + +ClientSocket::~ClientSocket() { + ShutDown(); + executor_.Shutdown(); + NEARBY_LOGS(INFO) << "ClientSocket gone."; +} + +void ClientSocket::Connect() { + if (state_ != State::kStateDisconnected) { + return; + } + + state_ = State::kStateClientConnectionRequest; + WriteControlPacket( + Packet::CreateConnectionRequestPacket(kProtocolVersion, kProtocolVersion, + GetConnection().GetMaxPacketSize(), + initial_data_provider_->Provide()) + .value()); + + // We are now waiting for server confirmation, which we should receive in + // OnReceiveControlPacket(). + state_ = State::kStateServerConfirm; +} + +void ClientSocket::DisconnectQuietly() { + BaseSocket::DisconnectQuietly(); + state_ = State::kStateDisconnected; +} + +void ClientSocket::OnReceiveControlPacket(Packet packet) { + if (packet.GetControlCommandNumber() == + Packet::ControlPacketType::kControlError) { + DisconnectQuietly(); + return; + } + + if (state_ != State::kStateServerConfirm) { + GetSocketCallback().on_error_cb(absl::InternalError( + absl::StrCat("unexpected control packet: ", packet.ToString()))); + return; + } + if (packet.GetControlCommandNumber() != + Packet::ControlPacketType::kControlConnectionConfirm) { + DisconnectInternal(absl::InternalError( + absl::StrCat("expected connection confirm packet but received ", + packet.ToString()))); + return; + } + if (packet.GetPayload().size() < kConnectionConfirmPacketMinLength) { + DisconnectInternal( + absl::InvalidArgumentError("packet of insufficient length received.")); + return; + } + int protocol_version = + ((packet.GetPayload().data()[0] << 8) | packet.GetPayload().data()[1]) & + (int16_t)0xFFFF; + if (protocol_version != kProtocolVersion) { + DisconnectInternal(absl::InternalError( + absl::StrCat("unexpected protocol version ", protocol_version))); + return; + } + int max_packet_size = + ((packet.GetPayload().data()[2] << 8) | packet.GetPayload().data()[3]) & + (int16_t)0xFFFF; + if (max_packet_size > GetConnection().GetMaxPacketSize()) { + DisconnectInternal(absl::InternalError( + absl::StrCat("server confirmed max packet size ", max_packet_size, + " higher than connection's max packet size ", + GetConnection().GetMaxPacketSize()))); + return; + } + state_ = State::kStateHandshakeCompleted; + OnConnected(max_packet_size); + + if (packet.GetPayload().size() > kConnectionConfirmPacketMinLength) { + std::string remaining_data = + packet.GetPayload().substr(kConnectionConfirmPacketMinLength); + GetSocketCallback().on_receive_cb(remaining_data); + } +} + +} // namespace weave +} // namespace nearby diff --git a/internal/weave/sockets/client_socket.h b/internal/weave/sockets/client_socket.h new file mode 100644 index 00000000..af21f0c3 --- /dev/null +++ b/internal/weave/sockets/client_socket.h @@ -0,0 +1,60 @@ +// 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_SOCKETS_CLIENT_SOCKET_H_ +#define THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKETS_CLIENT_SOCKET_H_ + +#include + +#include "internal/weave/base_socket.h" +#include "internal/weave/connection.h" +#include "internal/weave/socket_callback.h" +#include "internal/weave/sockets/initial_data_provider.h" + +namespace nearby { +namespace weave { + +// The Weave Client socket class. This socket implements handling of control +// packets in accordance with the Weave handshake. +class ClientSocket : public BaseSocket { + public: + ClientSocket(const Connection& connection, SocketCallback&& socket_callback); + // Note that when using this constructor, ClientSocket will take ownership of + // the InitialDataProvider instance. + ClientSocket(const Connection& connection, SocketCallback&& socket_callback, + std::unique_ptr initial_data_provider); + ~ClientSocket() override; + void Connect() override; + + protected: + void DisconnectQuietly() override; + void OnReceiveControlPacket(Packet packet) override; + + private: + static constexpr int kProtocolVersion = 0b0001; + enum class State { + kStateDisconnected = 0, + kStateClientConnectionRequest = 1, + kStateServerConfirm = 2, + kStateHandshakeCompleted = 3, + }; + std::unique_ptr initial_data_provider_; + SingleThreadExecutor executor_; + State state_ = State::kStateDisconnected; +}; + +} // namespace weave +} // namespace nearby + +#endif // THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKETS_CLIENT_SOCKET_H_ diff --git a/internal/weave/sockets/client_socket_test.cc b/internal/weave/sockets/client_socket_test.cc new file mode 100644 index 00000000..b3b994a3 --- /dev/null +++ b/internal/weave/sockets/client_socket_test.cc @@ -0,0 +1,365 @@ +// 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/sockets/client_socket.h" + +#include +#include +#include +#include +#include + +#include "gmock/gmock.h" +#include "protobuf-matchers/protocol-buffer-matchers.h" +#include "gtest/gtest.h" +#include "absl/time/clock.h" +#include "absl/time/time.h" +#include "internal/platform/byte_array.h" +#include "internal/platform/logging.h" +#include "internal/weave/sockets/initial_data_provider.h" + +namespace nearby { +namespace weave { +namespace { + +constexpr int kClientMaxPacketSize = 4; +constexpr int kServerMaxPacketSize = 3; +constexpr int kProtocolVersion = 1; + +class TestClientSocket : public ClientSocket { + public: + TestClientSocket(const Connection& connection, SocketCallback&& callback) + : ClientSocket(connection, std::move(callback)) {} + TestClientSocket(const Connection& connection, SocketCallback&& callback, + std::unique_ptr provider) + : ClientSocket(connection, std::move(callback), std::move(provider)) {} + void OnReceiveControlPacketProxy(Packet packet) { + OnReceiveControlPacket(std::move(packet)); + } +}; + +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() const 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 SetMaxPacketSize(int packet_size) { max_packet_size_ = packet_size; } + 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); + } + + protected: + int max_packet_size_; + ConnectionCallback callback_; + std::vector packets_written_; + bool instant_transmit_ = true; + bool open_ = false; +}; + +class ClientSocketTest : public ::testing::Test { + public: + ClientSocketTest() + : connection_(FakeConnection(kClientMaxPacketSize)), + socket_(TestClientSocket(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) { + last_error_ = status; + NEARBY_LOGS(ERROR) << status; + }, + })) {} + void SetUp() override { EXPECT_FALSE(socket_.IsConnected()); } + + void TearDown() override { + EXPECT_EQ(socket_.IsConnected(), expect_connected_); + } + void RunConnect(int client_size, int server_size, + absl::string_view initial_data) { + connection_.SetMaxPacketSize(client_size); + NEARBY_LOGS(INFO) << "connect"; + socket_.Connect(); + absl::SleepFor(absl::Milliseconds(10)); + auto packet = Packet::FromBytes(ByteArray(connection_.PollWrittenPacket())); + ASSERT_OK(packet); + EXPECT_TRUE(packet->IsControlPacket()); + EXPECT_EQ(packet->GetControlCommandNumber(), + Packet::ControlPacketType::kControlConnectionRequest); + EXPECT_EQ(packet->GetPacketCounter(), 0); + int16_t min_protocol_version = ((packet->GetPayload().data()[0] << 8) | + packet->GetPayload().data()[1]) & + (int16_t)0xFFFF; + int16_t max_protocol_version = ((packet->GetPayload().data()[2] << 8) | + packet->GetPayload().data()[3]) & + (int16_t)0xFFFF; + EXPECT_EQ(min_protocol_version, kProtocolVersion); + EXPECT_EQ(max_protocol_version, kProtocolVersion); + int pkt_size = ((packet->GetPayload().data()[4] << 8) | + packet->GetPayload().data()[5]) & + (int16_t)0xFFFF; + EXPECT_EQ(pkt_size, client_size); + EXPECT_EQ(packet->GetPayload().size(), 6); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionConfirmPacket(kProtocolVersion, server_size, + initial_data) + .value()); + int expected_payload_size = server_size - 1; + auto status = socket_.Write(ByteArray(expected_payload_size)); + EXPECT_OK(status.Get().GetResult()); + absl::SleepFor(absl::Milliseconds(10)); + ASSERT_FALSE(connection_.NoMorePackets()); + connection_.PollWrittenPacket(); + EXPECT_TRUE(connection_.NoMorePackets()); + status = socket_.Write(ByteArray(expected_payload_size + 1)); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_OK(status.Get().GetResult()); + ASSERT_FALSE(connection_.NoMorePackets()); + connection_.PollWrittenPacket(); + connection_.PollWrittenPacket(); + EXPECT_TRUE(connection_.NoMorePackets()); + } + FakeConnection connection_; + TestClientSocket socket_; + Mutex mutex_; + bool connected_ = true; + std::vector messages_read_; + bool expect_connected_ = true; + absl::Status last_error_; +}; + +TEST_F(ClientSocketTest, TestClientServerPacketSizeSame) { + RunConnect(2, 2, ""); +} + +TEST_F(ClientSocketTest, TestClientServerPacketSizeLower) { + RunConnect(kClientMaxPacketSize, kServerMaxPacketSize, ""); +} + +TEST_F(ClientSocketTest, TestClientServerPacketSizeHigher) { + int client_size = 4; + int server_size = 5; + connection_.SetMaxPacketSize(client_size); + socket_.Connect(); + absl::SleepFor(absl::Milliseconds(10)); + auto packet = Packet::FromBytes(ByteArray(connection_.PollWrittenPacket())); + ASSERT_OK(packet); + EXPECT_TRUE(packet->IsControlPacket()); + EXPECT_EQ(packet->GetControlCommandNumber(), + Packet::ControlPacketType::kControlConnectionRequest); + EXPECT_EQ(packet->GetPacketCounter(), 0); + int16_t min_protocol_version = + ((packet->GetPayload().data()[0] << 8) | packet->GetPayload().data()[1]) & + (int16_t)0xFFFF; + int16_t max_protocol_version = + ((packet->GetPayload().data()[2] << 8) | packet->GetPayload().data()[3]) & + (int16_t)0xFFFF; + EXPECT_EQ(min_protocol_version, kProtocolVersion); + EXPECT_EQ(max_protocol_version, kProtocolVersion); + int pkt_size = + ((packet->GetPayload().data()[4] << 8) | packet->GetPayload().data()[5]) & + (int16_t)0xFFFF; + EXPECT_EQ(pkt_size, client_size); + EXPECT_EQ(packet->GetPayload().size(), 6); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionConfirmPacket(kProtocolVersion, server_size, "") + .value()); + absl::SleepFor(absl::Milliseconds(10)); + expect_connected_ = false; +} + +TEST_F(ClientSocketTest, TestConnectInitialData) { + std::string initial_data = "12"; + RunConnect(2, 2, initial_data); + EXPECT_EQ(messages_read_.size(), 1); + EXPECT_EQ(messages_read_[0], initial_data); +} + +TEST_F(ClientSocketTest, TestConnectConnect) { + RunConnect(2, 2, ""); + ASSERT_TRUE(connection_.NoMorePackets()); + socket_.Connect(); + absl::SleepFor(absl::Milliseconds(10)); + EXPECT_TRUE(connection_.NoMorePackets()); +} + +TEST_F(ClientSocketTest, TestResponseBeforeOnTransmit) { + connection_.SetInstantTransmit(false); + socket_.Connect(); + absl::SleepFor(absl::Milliseconds(10)); + ASSERT_FALSE(connection_.NoMorePackets()); + auto packet = Packet::FromBytes(ByteArray(connection_.PollWrittenPacket())); + EXPECT_TRUE(packet->IsControlPacket()); + EXPECT_EQ(packet->GetControlCommandNumber(), + Packet::ControlPacketType::kControlConnectionRequest); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionConfirmPacket(kProtocolVersion, 2, "").value()); + absl::SleepFor(absl::Milliseconds(10)); + connection_.OnTransmitProxy(absl::OkStatus()); +} + +TEST_F(ClientSocketTest, TestDisconnect) { + RunConnect(2, 2, ""); + ASSERT_TRUE(connection_.NoMorePackets()); + socket_.Disconnect(); + absl::SleepFor(absl::Milliseconds(10)); + ASSERT_FALSE(connection_.NoMorePackets()); + Packet expected = Packet::CreateErrorPacket(); + EXPECT_OK(expected.SetPacketCounter(4)); + EXPECT_EQ(connection_.PollWrittenPacket(), expected.GetBytes()); + expect_connected_ = false; +} + +TEST_F(ClientSocketTest, TestDisconnectConnect) { + RunConnect(2, 2, ""); + ASSERT_TRUE(connection_.NoMorePackets()); + socket_.Disconnect(); + absl::SleepFor(absl::Milliseconds(10)); + ASSERT_FALSE(connection_.NoMorePackets()); + Packet expected = Packet::CreateErrorPacket(); + EXPECT_OK(expected.SetPacketCounter(4)); + EXPECT_EQ(connection_.PollWrittenPacket(), expected.GetBytes()); + RunConnect(2, 2, ""); + absl::SleepFor(absl::Milliseconds(10)); + ASSERT_TRUE(connection_.NoMorePackets()); + expect_connected_ = true; +} + +TEST_F(ClientSocketTest, TestReceiveErrorPacket) { + RunConnect(2, 2, ""); + ASSERT_TRUE(connection_.NoMorePackets()); + socket_.OnReceiveControlPacketProxy(Packet::CreateErrorPacket()); + absl::SleepFor(absl::Milliseconds(10)); + expect_connected_ = false; +} + +TEST_F(ClientSocketTest, TestNoConnectionConfirmWrongCommandNumber) { + socket_.Connect(); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionRequestPacket(0, 0, 0, "").value()); + expect_connected_ = false; +} + +TEST_F(ClientSocketTest, TestUnexpectedControlPacketNoConnection) { + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionRequestPacket(0, 0, 0, "").value()); + expect_connected_ = false; + EXPECT_EQ(last_error_.code(), absl::StatusCode::kInternal); +} + +TEST_F(ClientSocketTest, TestUnexpectedControlPacketConnected) { + RunConnect(2, 2, ""); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionRequestPacket(0, 0, 0, "").value()); + EXPECT_EQ(last_error_.code(), absl::StatusCode::kInternal); +} + +TEST_F(ClientSocketTest, TestBadConnectionConfirmPacket) { + socket_.Connect(); + auto packet = + Packet::CreateConnectionConfirmPacket(kProtocolVersion, 6, "").value(); + socket_.OnReceiveControlPacketProxy( + *Packet::FromBytes(ByteArray(packet.GetBytes().substr(0, 2)))); + expect_connected_ = false; + EXPECT_EQ(last_error_.code(), absl::StatusCode::kInvalidArgument); +} + +TEST_F(ClientSocketTest, TestUnsupportedProtocolVersion) { + socket_.Connect(); + socket_.OnReceiveControlPacketProxy( + Packet::CreateConnectionConfirmPacket(kProtocolVersion + 1, + kServerMaxPacketSize, "") + .value()); + expect_connected_ = false; +} + +TEST_F(ClientSocketTest, TestSocketWithRandomDataProvider) { + auto provider = std::make_unique(/*number_of_bytes=*/5); + TestClientSocket 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 = + [](absl::Status status) { NEARBY_LOGS(ERROR) << status; }, + }, + std::move(provider)); + socket.Connect(); + absl::SleepFor(absl::Milliseconds(10)); + Packet packet = + *Packet::FromBytes(ByteArray(connection_.PollWrittenPacket())); + EXPECT_TRUE(packet.IsControlPacket()); + EXPECT_EQ(packet.GetControlCommandNumber(), + Packet::ControlPacketType::kControlConnectionRequest); + // 2 + 2 + 2 + 5. + EXPECT_EQ(packet.GetPayload().size(), 11); + EXPECT_NE(packet.GetPayload().substr(5), "\x00\x00\x00\x00\x00"); + // we never sent back a connection confirm so we are still disconnected. + expect_connected_ = false; +} + +} // namespace +} // namespace weave +} // namespace nearby diff --git a/internal/weave/sockets/initial_data_provider.h b/internal/weave/sockets/initial_data_provider.h new file mode 100644 index 00000000..cef403a3 --- /dev/null +++ b/internal/weave/sockets/initial_data_provider.h @@ -0,0 +1,49 @@ +// 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_SOCKETS_INITIAL_DATA_PROVIDER_H_ +#define THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKETS_INITIAL_DATA_PROVIDER_H_ + +#include + +#include "absl/random/random.h" + +namespace nearby { +namespace weave { + +class InitialDataProvider { + public: + virtual ~InitialDataProvider() = default; + virtual std::string Provide() { return ""; } +}; + +class RandomDataProvider : public InitialDataProvider { + public: + explicit RandomDataProvider(int number_of_bytes) + : number_of_bytes_(number_of_bytes) {} + std::string Provide() override { + uint64_t res = absl::uniform_int_distribution( + 1, pow(10, static_cast(number_of_bytes_)))(prng_); + return absl::StrCat(res); + } + + private: + int number_of_bytes_; + absl::BitGen prng_; +}; + +} // namespace weave +} // namespace nearby + +#endif // THIRD_PARTY_NEARBY_INTERNAL_WEAVE_SOCKETS_INITIAL_DATA_PROVIDER_H_