TrinityCore
Loading...
Searching...
No Matches
Socket.h
Go to the documentation of this file.
1/*
2 * This file is part of the TrinityCore Project. See AUTHORS file for Copyright information
3 *
4 * This program is free software; you can redistribute it and/or modify it
5 * under the terms of the GNU General Public License as published by the
6 * Free Software Foundation; either version 2 of the License, or (at your
7 * option) any later version.
8 *
9 * This program is distributed in the hope that it will be useful, but WITHOUT
10 * ANY WARRANTY; without even the implied warranty of MERCHANTABILITY or
11 * FITNESS FOR A PARTICULAR PURPOSE. See the GNU General Public License for
12 * more details.
13 *
14 * You should have received a copy of the GNU General Public License along
15 * with this program. If not, see <http://www.gnu.org/licenses/>.
16 */
17
18#ifndef TRINITYCORE_SOCKET_H
19#define TRINITYCORE_SOCKET_H
20
21#include "Concepts.h"
22#include "IoContext.h"
23#include "IpAddress.h"
24#include "Log.h"
25#include "MessageBuffer.h"
27#include <boost/asio/compose.hpp>
28#include <boost/asio/ip/tcp.hpp>
29#include <atomic>
30#include <memory>
31#include <queue>
32#include <type_traits>
33
34#ifdef BOOST_ASIO_HAS_IOCP
35#define TC_SOCKET_USE_IOCP
36#endif
37
38namespace Trinity::Net
39{
40using IoContextTcpSocket = boost::asio::basic_stream_socket<boost::asio::ip::tcp, Asio::IoContextExecutor>;
41
42namespace Impl::Operations
43{
44template <typename Socket>
45struct Connect;
46}
47
49{
51 Stop
52};
53
54inline boost::asio::mutable_buffer PrepareReadBuffer(MessageBuffer& readBuffer)
55{
56 readBuffer.Normalize();
57 readBuffer.EnsureFreeSpace();
58 return boost::asio::buffer(readBuffer.GetWritePointer(), readBuffer.GetRemainingSpace());
59}
60
61template <typename SocketType>
63{
65 {
66 return this->Socket->ReadHandler();
67 }
68
69 SocketType* Socket;
70};
71
72template <typename AsyncReadObjectType, typename ReadHandlerObjectType = AsyncReadObjectType>
74{
75 explicit ReadConnectionInitializer(AsyncReadObjectType* socket) : Socket(socket), ReadCallback({ .Socket = socket }) { }
76 explicit ReadConnectionInitializer(AsyncReadObjectType* socket, ReadHandlerObjectType* callbackSocket) : Socket(socket), ReadCallback({ .Socket = callbackSocket }) { }
77
78 void Start() override
79 {
80 Socket->AsyncRead(std::move(ReadCallback));
81
82 this->InvokeNext();
83 }
84
85 AsyncReadObjectType* Socket;
87};
88
125template<class Stream = IoContextTcpSocket>
126class Socket : public std::enable_shared_from_this<Socket<Stream>>
127{
128public:
129 template<typename... Args>
130 explicit Socket(IoContextTcpSocket&& socket, Args&&... args) : _socket(std::move(socket), std::forward<Args>(args)...),
132 {
133 }
134
135 template<typename... Args>
136 explicit Socket(Asio::IoContext& context, Args&&... args) : _socket(context, std::forward<Args>(args)...),
138 {
139 }
140
141 Socket(Socket const& other) = delete;
142 Socket(Socket&& other) = delete;
143 Socket& operator=(Socket const& other) = delete;
144 Socket& operator=(Socket&& other) = delete;
145
146 virtual ~Socket()
147 {
149 boost::system::error_code error;
150 _socket.close(error);
151 }
152
153 virtual void Start() { }
154
155 template <BOOST_ASIO_COMPLETION_TOKEN_FOR(void(boost::system::error_code, boost::asio::ip::tcp::endpoint)) Callback>
156 decltype(auto) Connect(boost::asio::ip::tcp::endpoint const& endpoint, Callback&& callback)
157 {
159 return boost::asio::async_compose<Callback, void(boost::system::error_code, boost::asio::ip::tcp::endpoint), Impl::Operations::Connect<Socket>>(
160 Impl::Operations::Connect<Socket>(this->shared_from_this(), endpoint), callback, this->underlying_stream());
161 }
162
163 template <BOOST_ASIO_COMPLETION_TOKEN_FOR(void(boost::system::error_code, boost::asio::ip::tcp::endpoint)) Callback>
164 decltype(auto) Connect(std::vector<boost::asio::ip::tcp::endpoint> const& endpoints, Callback&& callback)
165 {
167 return boost::asio::async_compose<Callback, void(boost::system::error_code, boost::asio::ip::tcp::endpoint), Impl::Operations::Connect<Socket>>(
168 Impl::Operations::Connect<Socket>(this->shared_from_this(), endpoints), callback, this->underlying_stream());
169 }
170
171 virtual bool Update()
172 {
174 return false;
175
176#ifndef TC_SOCKET_USE_IOCP
178 return true;
179
180 for (; HandleQueue();)
181 ;
182#endif
183
184 return true;
185 }
186
187 boost::asio::ip::address const& GetRemoteIpAddress() const
188 {
190 }
191
193 {
194 return _remoteEndpoint.Port;
195 }
196
197 void SetRemoteEndpoint(boost::asio::ip::tcp::endpoint const& endpoint)
198 {
199 _remoteEndpoint = endpoint;
200 }
201
202 template <invocable_r<SocketReadCallbackResult> Callback>
203 void AsyncRead(Callback&& callback)
204 {
205 if (!IsOpen())
206 return;
207
208 _socket.async_read_some(PrepareReadBuffer(_readBuffer),
209 [self = this->shared_from_this(), callback = std::forward<Callback>(callback)](boost::system::error_code const& error, size_t transferredBytes) mutable
210 {
211 if (self->ReadHandlerInternal(error, transferredBytes))
213 self->AsyncRead(std::forward<Callback>(callback));
214 });
215 }
216
218 {
219 _writeQueue.push(std::move(buffer));
220
221#ifdef TC_SOCKET_USE_IOCP
223#endif
224 }
225
226 bool IsOpen() const { return _openState == OpenState_Open; }
227
229 {
231 return;
232
233 boost::system::error_code shutdownError;
234 _socket.shutdown(boost::asio::socket_base::shutdown_send, shutdownError);
235 if (shutdownError)
236 TC_LOG_DEBUG("network", "Socket::CloseSocket: {} errored when shutting down socket: {} ({})", GetRemoteIpAddress(),
237 shutdownError.value(), shutdownError.message());
238
239 this->OnClose();
240 }
241
244 {
245 uint8 oldState = OpenState_Open;
246 if (!_openState.compare_exchange_strong(oldState, OpenState_Closing))
247 return;
248
249 if (_writeQueue.empty())
250 CloseSocket();
251 }
252
254
256 {
257 return _socket;
258 }
259
260protected:
261 virtual void OnClose() { }
262
264
266 {
267 if (_isWritingAsync)
268 return false;
269
270 _isWritingAsync = true;
271
272#ifdef TC_SOCKET_USE_IOCP
273 MessageBuffer& buffer = _writeQueue.front();
274 _socket.async_write_some(boost::asio::buffer(buffer.GetReadPointer(), buffer.GetActiveSize()),
275 [self = this->shared_from_this()](boost::system::error_code const& error, std::size_t transferedBytes)
276 {
277 self->WriteHandler(error, transferedBytes);
278 });
279#else
280 _socket.async_wait(boost::asio::socket_base::wait_type::wait_write,
281 [self = this->shared_from_this()](boost::system::error_code const& error)
282 {
283 self->WriteHandlerWrapper(error);
284 });
285#endif
286
287 return false;
288 }
289
290 void SetNoDelay(bool enable)
291 {
292 boost::system::error_code err;
293 _socket.set_option(boost::asio::ip::tcp::no_delay(enable), err);
294 if (err)
295 TC_LOG_DEBUG("network", "Socket::SetNoDelay: failed to set_option(boost::asio::ip::tcp::no_delay) for {} - {} ({})",
296 GetRemoteIpAddress(), err.value(), err.message());
297 }
298
299private:
300 bool ReadHandlerInternal(boost::system::error_code const& error, size_t transferredBytes)
301 {
302 if (error)
303 {
304 CloseSocket();
305 return false;
306 }
307
308 _readBuffer.WriteCompleted(transferredBytes);
309 return IsOpen();
310 }
311
313 {
314 _writeQueue.pop();
315 if (_openState == OpenState_Closing && _writeQueue.empty())
316 CloseSocket();
317 }
318
319#ifdef TC_SOCKET_USE_IOCP
320
321 void WriteHandler(boost::system::error_code const& error, std::size_t transferedBytes)
322 {
323 if (!error)
324 {
325 _isWritingAsync = false;
326 _writeQueue.front().ReadCompleted(transferedBytes);
327 if (!_writeQueue.front().GetActiveSize())
329
330 if (!_writeQueue.empty())
332 }
333 else
334 CloseSocket();
335 }
336
337#else
338
339 void WriteHandlerWrapper(boost::system::error_code const& /*error*/)
340 {
341 _isWritingAsync = false;
342 HandleQueue();
343 }
344
346 {
347 if (_writeQueue.empty())
348 return false;
349
350 MessageBuffer& queuedMessage = _writeQueue.front();
351
352 std::size_t bytesToSend = queuedMessage.GetActiveSize();
353
354 boost::system::error_code error;
355 std::size_t bytesSent = _socket.write_some(boost::asio::buffer(queuedMessage.GetReadPointer(), bytesToSend), error);
356
357 if (error)
358 {
359 if (error == boost::asio::error::would_block || error == boost::asio::error::try_again)
360 return AsyncProcessQueue();
361
363 return false;
364 }
365 else if (bytesSent == 0)
366 {
368 return false;
369 }
370 else if (bytesSent < bytesToSend) // now n > 0
371 {
372 queuedMessage.ReadCompleted(bytesSent);
373 return AsyncProcessQueue();
374 }
375
377 return !_writeQueue.empty();
378 }
379
380#endif
381
382 Stream _socket;
383
384 struct Endpoint
385 {
386 Endpoint() : Address(), Port(0) { }
387 explicit(false) Endpoint(boost::asio::ip::tcp_endpoint const& endpoint) : Address(endpoint.address()), Port(endpoint.port()) { }
388
389 boost::asio::ip::address Address;
392
394 std::queue<MessageBuffer> _writeQueue;
395
396 // Socket open state "enum" (not enum to enable integral std::atomic api)
397 static constexpr uint8 OpenState_Open = 0x0;
398 static constexpr uint8 OpenState_Closing = 0x1;
399 static constexpr uint8 OpenState_Closed = 0x2;
400
401 std::atomic<uint8> _openState;
402
403 bool _isWritingAsync = false;
404};
405
406namespace Impl::Operations
407{
409{
410 explicit ConnectState(std::shared_ptr<void> const& socketRef, boost::asio::ip::tcp::endpoint const& endpoint)
411 : SocketRef(socketRef), Endpoints(1, endpoint), Index(-1) { }
412
413 explicit ConnectState(std::shared_ptr<void> const& socketRef, std::vector<boost::asio::ip::tcp::endpoint> const& endpoints)
414 : SocketRef(socketRef), Endpoints(endpoints), Index(-1) { }
415
416 std::weak_ptr<void> SocketRef;
417 std::vector<boost::asio::ip::tcp::endpoint> Endpoints;
418 std::ptrdiff_t Index;
419};
420
421template <typename Socket>
423{
424 explicit Connect(std::shared_ptr<Socket> const& socketRef, boost::asio::ip::tcp::endpoint const& endpoint)
425 : State(std::make_shared<ConnectState>(std::move(socketRef), endpoint)) { }
426
427 explicit Connect(std::shared_ptr<Socket> const& socketRef, std::vector<boost::asio::ip::tcp::endpoint> const& endpoints)
428 : State(std::make_shared<ConnectState>(std::move(socketRef), endpoints)) { }
429
430 std::shared_ptr<ConnectState> State;
431
432 template <typename Handler>
433 void operator()(Handler& handler, boost::system::error_code error = {})
434 {
435 std::shared_ptr<Socket> socket = static_pointer_cast<Socket>(State->SocketRef.lock());
436 if (!socket)
437 {
438 error = boost::asio::error::operation_aborted;
439 handler.complete(error, boost::asio::ip::tcp::endpoint());
440 return;
441 }
442
443 bool isFirst = State->Index < 0;
444
445 if (std::max(State->Index, std::ptrdiff_t(0)) >= std::ssize(State->Endpoints))
446 {
447 Connect::HandleError(socket.get(), "failed to connect to any of specified endpoints");
448 error = boost::asio::error::not_found;
449 handler.complete(error, boost::asio::ip::tcp::endpoint());
450 return;
451 }
452
453 if (!isFirst && !socket->underlying_stream().is_open())
454 {
455 Connect::HandleError(socket.get(), "socket closed");
456 error = boost::asio::error::operation_aborted;
457 handler.complete(error, boost::asio::ip::tcp::endpoint());
458 return;
459 }
460
461 if (!error && !isFirst)
462 {
463 socket->SetRemoteEndpoint(State->Endpoints[State->Index]);
464 handler.complete(error, State->Endpoints[State->Index]);
465 }
466 else
467 {
468#if BOOST_VERSION >= 107700
469 if (handler.get_cancellation_state().cancelled() != boost::asio::cancellation_type::none)
470 {
471 Connect::HandleError(socket.get(), "connect cancelled");
472 error = boost::asio::error::operation_aborted;
473 handler.complete(error, boost::asio::ip::tcp::endpoint());
474 return;
475 }
476#endif
477
478 socket->underlying_stream().close(error);
479 socket->underlying_stream().async_connect(State->Endpoints[++State->Index], std::move(handler));
480 }
481 }
482
483 static void HandleError(Socket* self, std::string_view message)
484 {
485 TC_LOG_DEBUG("network", "Socket::Connect: {}", message);
486 self->CloseSocket();
487 }
488};
489}
490}
491
492#endif // TRINITYCORE_SOCKET_H
uint8_t uint8
Definition Define.h:156
uint16_t uint16
Definition Define.h:155
#define TC_LOG_DEBUG(filterType__, message__,...)
Definition Log.h:173
size_type GetRemainingSpace() const
void ReadCompleted(size_type bytes)
void WriteCompleted(size_type bytes)
uint8 * GetReadPointer()
size_type GetActiveSize() const
uint8 * GetWritePointer()
void EnsureFreeSpace()
uint16 GetRemotePort() const
Definition Socket.h:192
decltype(auto) Connect(std::vector< boost::asio::ip::tcp::endpoint > const &endpoints, Callback &&callback)
Definition Socket.h:164
static constexpr uint8 OpenState_Closed
Definition Socket.h:399
std::atomic< uint8 > _openState
Definition Socket.h:401
void QueuedBufferWriteDone()
Definition Socket.h:312
Socket(Socket const &other)=delete
bool ReadHandlerInternal(boost::system::error_code const &error, size_t transferredBytes)
Definition Socket.h:300
void SetNoDelay(bool enable)
Definition Socket.h:290
Socket(IoContextTcpSocket &&socket, Args &&... args)
Definition Socket.h:130
std::queue< MessageBuffer > _writeQueue
Definition Socket.h:394
virtual SocketReadCallbackResult ReadHandler()
Definition Socket.h:263
Socket & operator=(Socket const &other)=delete
bool AsyncProcessQueue()
Definition Socket.h:265
bool IsOpen() const
Definition Socket.h:226
Stream & underlying_stream()
Definition Socket.h:255
boost::asio::ip::address const & GetRemoteIpAddress() const
Definition Socket.h:187
static constexpr uint8 OpenState_Closing
Transition to Closed state after sending all queued data.
Definition Socket.h:398
static constexpr uint8 OpenState_Open
Definition Socket.h:397
virtual void OnClose()
Definition Socket.h:261
virtual bool Update()
Definition Socket.h:171
void AsyncRead(Callback &&callback)
Definition Socket.h:203
void QueuePacket(MessageBuffer &&buffer)
Definition Socket.h:217
void DelayedCloseSocket()
Marks the socket for closing after write buffer becomes empty.
Definition Socket.h:243
MessageBuffer _readBuffer
Definition Socket.h:393
void WriteHandlerWrapper(boost::system::error_code const &)
Definition Socket.h:339
Socket(Asio::IoContext &context, Args &&... args)
Definition Socket.h:136
virtual void Start()
Definition Socket.h:153
decltype(auto) Connect(boost::asio::ip::tcp::endpoint const &endpoint, Callback &&callback)
Definition Socket.h:156
void SetRemoteEndpoint(boost::asio::ip::tcp::endpoint const &endpoint)
Definition Socket.h:197
Socket(Socket &&other)=delete
Socket & operator=(Socket &&other)=delete
struct Trinity::Net::Socket::Endpoint _remoteEndpoint
virtual ~Socket()
Definition Socket.h:146
MessageBuffer & GetReadBuffer()
Definition Socket.h:253
SocketReadCallbackResult
Definition Socket.h:49
boost::asio::basic_stream_socket< boost::asio::ip::tcp, Asio::IoContextExecutor > IoContextTcpSocket
Definition Socket.h:40
boost::asio::mutable_buffer PrepareReadBuffer(MessageBuffer &readBuffer)
Definition Socket.h:54
STL namespace.
ConnectState(std::shared_ptr< void > const &socketRef, boost::asio::ip::tcp::endpoint const &endpoint)
Definition Socket.h:410
ConnectState(std::shared_ptr< void > const &socketRef, std::vector< boost::asio::ip::tcp::endpoint > const &endpoints)
Definition Socket.h:413
std::vector< boost::asio::ip::tcp::endpoint > Endpoints
Definition Socket.h:417
void operator()(Handler &handler, boost::system::error_code error={})
Definition Socket.h:433
std::shared_ptr< ConnectState > State
Definition Socket.h:430
Connect(std::shared_ptr< Socket > const &socketRef, boost::asio::ip::tcp::endpoint const &endpoint)
Definition Socket.h:424
static void HandleError(Socket *self, std::string_view message)
Definition Socket.h:483
Connect(std::shared_ptr< Socket > const &socketRef, std::vector< boost::asio::ip::tcp::endpoint > const &endpoints)
Definition Socket.h:427
SocketReadCallbackResult operator()() const
Definition Socket.h:64
AsyncReadObjectType * Socket
Definition Socket.h:85
InvokeReadHandlerCallback< ReadHandlerObjectType > ReadCallback
Definition Socket.h:86
ReadConnectionInitializer(AsyncReadObjectType *socket, ReadHandlerObjectType *callbackSocket)
Definition Socket.h:76
ReadConnectionInitializer(AsyncReadObjectType *socket)
Definition Socket.h:75
boost::asio::ip::address Address
Definition Socket.h:389