1 //===-- SocketTest.cpp ------------------------------------------*- C++ -*-===//
2 //
3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4 // See https://llvm.org/LICENSE.txt for license information.
5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6 //
7 //===----------------------------------------------------------------------===//
8 
9 #include <cstdio>
10 #include <functional>
11 #include <thread>
12 
13 #include "gtest/gtest.h"
14 
15 #include "lldb/Host/Config.h"
16 #include "lldb/Host/Socket.h"
17 #include "lldb/Host/common/TCPSocket.h"
18 #include "lldb/Host/common/UDPSocket.h"
19 #include "llvm/Support/FileSystem.h"
20 #include "llvm/Support/Path.h"
21 #include "llvm/Testing/Support/Error.h"
22 
23 #ifndef LLDB_DISABLE_POSIX
24 #include "lldb/Host/posix/DomainSocket.h"
25 #endif
26 
27 using namespace lldb_private;
28 
29 class SocketTest : public testing::Test {
30 public:
31   void SetUp() override {
32     ASSERT_THAT_ERROR(Socket::Initialize(), llvm::Succeeded());
33   }
34 
35   void TearDown() override { Socket::Terminate(); }
36 
37 protected:
38   static void AcceptThread(Socket *listen_socket,
39                            bool child_processes_inherit, Socket **accept_socket,
40                            Status *error) {
41     *error = listen_socket->Accept(*accept_socket);
42   }
43 
44   template <typename SocketType>
45   void CreateConnectedSockets(
46       llvm::StringRef listen_remote_address,
47       const std::function<std::string(const SocketType &)> &get_connect_addr,
48       std::unique_ptr<SocketType> *a_up, std::unique_ptr<SocketType> *b_up) {
49     bool child_processes_inherit = false;
50     Status error;
51     std::unique_ptr<SocketType> listen_socket_up(
52         new SocketType(true, child_processes_inherit));
53     EXPECT_FALSE(error.Fail());
54     error = listen_socket_up->Listen(listen_remote_address, 5);
55     EXPECT_FALSE(error.Fail());
56     EXPECT_TRUE(listen_socket_up->IsValid());
57 
58     Status accept_error;
59     Socket *accept_socket;
60     std::thread accept_thread(AcceptThread, listen_socket_up.get(),
61                               child_processes_inherit, &accept_socket,
62                               &accept_error);
63 
64     std::string connect_remote_address = get_connect_addr(*listen_socket_up);
65     std::unique_ptr<SocketType> connect_socket_up(
66         new SocketType(true, child_processes_inherit));
67     EXPECT_FALSE(error.Fail());
68     error = connect_socket_up->Connect(connect_remote_address);
69     EXPECT_FALSE(error.Fail());
70     EXPECT_TRUE(connect_socket_up->IsValid());
71 
72     a_up->swap(connect_socket_up);
73     EXPECT_TRUE(error.Success());
74     EXPECT_NE(nullptr, a_up->get());
75     EXPECT_TRUE((*a_up)->IsValid());
76 
77     accept_thread.join();
78     b_up->reset(static_cast<SocketType *>(accept_socket));
79     EXPECT_TRUE(accept_error.Success());
80     EXPECT_NE(nullptr, b_up->get());
81     EXPECT_TRUE((*b_up)->IsValid());
82 
83     listen_socket_up.reset();
84   }
85 };
86 
87 TEST_F(SocketTest, DecodeHostAndPort) {
88   std::string host_str;
89   std::string port_str;
90   int32_t port;
91   Status error;
92   EXPECT_TRUE(Socket::DecodeHostAndPort("localhost:1138", host_str, port_str,
93                                         port, &error));
94   EXPECT_STREQ("localhost", host_str.c_str());
95   EXPECT_STREQ("1138", port_str.c_str());
96   EXPECT_EQ(1138, port);
97   EXPECT_TRUE(error.Success());
98 
99   EXPECT_FALSE(Socket::DecodeHostAndPort("google.com:65536", host_str, port_str,
100                                          port, &error));
101   EXPECT_TRUE(error.Fail());
102   EXPECT_STREQ("invalid host:port specification: 'google.com:65536'",
103                error.AsCString());
104 
105   EXPECT_FALSE(Socket::DecodeHostAndPort("google.com:-1138", host_str, port_str,
106                                          port, &error));
107   EXPECT_TRUE(error.Fail());
108   EXPECT_STREQ("invalid host:port specification: 'google.com:-1138'",
109                error.AsCString());
110 
111   EXPECT_FALSE(Socket::DecodeHostAndPort("google.com:65536", host_str, port_str,
112                                          port, &error));
113   EXPECT_TRUE(error.Fail());
114   EXPECT_STREQ("invalid host:port specification: 'google.com:65536'",
115                error.AsCString());
116 
117   EXPECT_TRUE(
118       Socket::DecodeHostAndPort("12345", host_str, port_str, port, &error));
119   EXPECT_STREQ("", host_str.c_str());
120   EXPECT_STREQ("12345", port_str.c_str());
121   EXPECT_EQ(12345, port);
122   EXPECT_TRUE(error.Success());
123 
124   EXPECT_TRUE(
125       Socket::DecodeHostAndPort("*:0", host_str, port_str, port, &error));
126   EXPECT_STREQ("*", host_str.c_str());
127   EXPECT_STREQ("0", port_str.c_str());
128   EXPECT_EQ(0, port);
129   EXPECT_TRUE(error.Success());
130 
131   EXPECT_TRUE(
132       Socket::DecodeHostAndPort("*:65535", host_str, port_str, port, &error));
133   EXPECT_STREQ("*", host_str.c_str());
134   EXPECT_STREQ("65535", port_str.c_str());
135   EXPECT_EQ(65535, port);
136   EXPECT_TRUE(error.Success());
137 
138   EXPECT_TRUE(
139       Socket::DecodeHostAndPort("[::1]:12345", host_str, port_str, port, &error));
140   EXPECT_STREQ("::1", host_str.c_str());
141   EXPECT_STREQ("12345", port_str.c_str());
142   EXPECT_EQ(12345, port);
143   EXPECT_TRUE(error.Success());
144 
145   EXPECT_TRUE(
146       Socket::DecodeHostAndPort("[abcd:12fg:AF58::1]:12345", host_str, port_str, port, &error));
147   EXPECT_STREQ("abcd:12fg:AF58::1", host_str.c_str());
148   EXPECT_STREQ("12345", port_str.c_str());
149   EXPECT_EQ(12345, port);
150   EXPECT_TRUE(error.Success());
151 }
152 
153 #ifndef LLDB_DISABLE_POSIX
154 TEST_F(SocketTest, DomainListenConnectAccept) {
155   llvm::SmallString<64> Path;
156   std::error_code EC = llvm::sys::fs::createUniqueDirectory("DomainListenConnectAccept", Path);
157   ASSERT_FALSE(EC);
158   llvm::sys::path::append(Path, "test");
159 
160   std::unique_ptr<DomainSocket> socket_a_up;
161   std::unique_ptr<DomainSocket> socket_b_up;
162   CreateConnectedSockets<DomainSocket>(
163       Path, [=](const DomainSocket &) { return Path.str().str(); },
164       &socket_a_up, &socket_b_up);
165 }
166 #endif
167 
168 TEST_F(SocketTest, TCPListen0ConnectAccept) {
169   std::unique_ptr<TCPSocket> socket_a_up;
170   std::unique_ptr<TCPSocket> socket_b_up;
171   CreateConnectedSockets<TCPSocket>(
172       "127.0.0.1:0",
173       [=](const TCPSocket &s) {
174         char connect_remote_address[64];
175         snprintf(connect_remote_address, sizeof(connect_remote_address),
176                  "127.0.0.1:%u", s.GetLocalPortNumber());
177         return std::string(connect_remote_address);
178       },
179       &socket_a_up, &socket_b_up);
180 }
181 
182 TEST_F(SocketTest, TCPGetAddress) {
183   std::unique_ptr<TCPSocket> socket_a_up;
184   std::unique_ptr<TCPSocket> socket_b_up;
185   CreateConnectedSockets<TCPSocket>(
186       "127.0.0.1:0",
187       [=](const TCPSocket &s) {
188         char connect_remote_address[64];
189         snprintf(connect_remote_address, sizeof(connect_remote_address),
190                  "127.0.0.1:%u", s.GetLocalPortNumber());
191         return std::string(connect_remote_address);
192       },
193       &socket_a_up, &socket_b_up);
194 
195   EXPECT_EQ(socket_a_up->GetLocalPortNumber(),
196             socket_b_up->GetRemotePortNumber());
197   EXPECT_EQ(socket_b_up->GetLocalPortNumber(),
198             socket_a_up->GetRemotePortNumber());
199   EXPECT_NE(socket_a_up->GetLocalPortNumber(),
200             socket_b_up->GetLocalPortNumber());
201   EXPECT_STREQ("127.0.0.1", socket_a_up->GetRemoteIPAddress().c_str());
202   EXPECT_STREQ("127.0.0.1", socket_b_up->GetRemoteIPAddress().c_str());
203 }
204 
205 TEST_F(SocketTest, UDPConnect) {
206   Socket *socket;
207 
208   bool child_processes_inherit = false;
209   auto error = UDPSocket::Connect("127.0.0.1:0", child_processes_inherit,
210                                   socket);
211 
212   std::unique_ptr<Socket> socket_up(socket);
213 
214   EXPECT_TRUE(error.Success());
215   EXPECT_TRUE(socket_up->IsValid());
216 }
217 
218 TEST_F(SocketTest, TCPListen0GetPort) {
219   Socket *server_socket;
220   Predicate<uint16_t> port_predicate;
221   port_predicate.SetValue(0, eBroadcastNever);
222   Status err =
223       Socket::TcpListen("10.10.12.3:0", false, server_socket, &port_predicate);
224   std::unique_ptr<TCPSocket> socket_up((TCPSocket*)server_socket);
225   EXPECT_TRUE(socket_up->IsValid());
226   EXPECT_NE(socket_up->GetLocalPortNumber(), 0);
227 }
228