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