TrinityCore
Loading...
Searching...
No Matches
Session.cpp
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#include "Session.h"
19#include "ByteConverter.h"
20#include "DatabaseEnv.h"
21#include "Errors.h"
23#include "MapUtils.h"
24#include "Memory.h"
25#include "QueryCallback.h"
26#include "ServiceDispatcher.h"
27#include "SessionManager.h"
28#include "SslContext.h"
29#include "rpc_types.pb.h"
30
32{
33 // ba.id, ba.email, ba.locked, ba.lock_country, ba.last_ip, ba.LoginTicketExpiry, bab.unbandate > UNIX_TIMESTAMP() OR bab.unbandate = bab.bandate, bab.unbandate = bab.bandate FROM battlenet_accounts ba LEFT JOIN battlenet_account_bans bab WHERE email = ?
34 Field* fields = result->Fetch();
35 Id = fields[0].GetUInt32();
36 Login = fields[1].GetString();
37 IsLockedToIP = fields[2].GetBool();
38 LockCountry = fields[3].GetString();
39 LastIP = fields[4].GetString();
40 LoginTicketExpiry = fields[5].GetUInt32();
41 IsBanned = fields[6].GetUInt64() != 0;
42 IsPermanenetlyBanned = fields[7].GetUInt64() != 0;
43
44 static constexpr uint32 GameAccountFieldsOffset = 8;
45
46 do
47 {
48 GameAccounts[result->Fetch()[GameAccountFieldsOffset].GetUInt32()].LoadResult(result->Fetch() + GameAccountFieldsOffset);
49
50 } while (result->NextRow());
51}
52
54{
55 // a.id, a.username, ab.bandate, ab.unbandate, ab.unbandate = ab.bandate, aa.SecurityLevel
56 Id = fields[0].GetUInt32();
57 Name = fields[1].GetString();
58 BanDate = fields[2].GetUInt32();
59 UnbanDate = fields[3].GetUInt32();
60 IsPermanenetlyBanned = fields[4].GetUInt32() != 0;
61 IsBanned = IsPermanenetlyBanned || UnbanDate > time(nullptr);
62 SecurityLevel = AccountTypes(fields[5].GetUInt8());
63
64 std::size_t hashPos = Name.find('#');
65 if (hashPos != std::string::npos)
66 DisplayName = std::string("WoW") + Name.substr(hashPos + 1);
67 else
68 DisplayName = Name;
69}
70
71constexpr std::size_t PacketHeaderLengthSize = sizeof(uint16);
72
73Battlenet::Session::Session(Trinity::Net::IoContextTcpSocket&& socket) : _socket(CreateSocket(std::move(socket))),
74 _packetReadState(PacketReadState::HeaderLength), _packetBuffer(PacketHeaderLengthSize),
75 _sessionId(++sSessionMgr.SessionIdGenerator), _creationTime(SystemTimePoint::clock::now()),
76 _accountInfo(new AccountInfo()), _gameAccountInfo(nullptr), _locale(),
77 _os(), _build(0), _buildVariant(), _timezoneOffset(0min), _ipCountry(), _clientSecret(), _authed(false), _requestToken(0)
78{
79}
80
82
83std::shared_ptr<Battlenet::Session::Socket> Battlenet::Session::CreateSocket(Trinity::Net::IoContextTcpSocket&& socket)
84{
85 return std::make_shared<Socket>(std::move(socket), SslContext::instance());
86}
87
89{
90 TC_LOG_TRACE("session", "{} Accepted connection", GetClientInfo());
91
92 // build initializer chain
93 std::array<std::shared_ptr<Trinity::Net::SocketConnectionInitializer>, 3> initializers =
94 { {
95 std::make_shared<Trinity::Net::IpBanCheckConnectionInitializer<Session>>(this),
96 std::make_shared<Trinity::Net::SslHandshakeConnectionInitializer<Socket>>(_socket.get()),
97 std::make_shared<Trinity::Net::ReadConnectionInitializer<Socket, Session>>(_socket.get(), this),
98 } };
99
101}
102
104{
105 if (!_socket->Update())
106 return false;
107
108 _queryProcessor.ProcessReadyCallbacks();
109
110 return true;
111}
112
114{
115 if (!_accountInfo)
116 return nullptr;
117
118 return Trinity::Containers::MapGetValuePtr(_accountInfo->GameAccounts, gameAccountId);
119}
120
122{
123 if (!_socket->IsOpen())
124 return;
125
126 _socket->QueuePacket(std::move(*packet));
127}
128
129void Battlenet::Session::SendResponse(uint32 token, pb::Message const* response)
130{
131 bgs::protocol::Header header;
132 header.set_token(token);
133 header.set_service_id(0xFE);
134 header.set_size(response->ByteSize());
135 header.set_allocated_ciid(&_clientInstanceId);
136
137 auto ciidGuard = Trinity::make_unique_ptr_with_deleter<&bgs::protocol::Header::release_ciid>(&header);
138
139 uint16 headerSize = header.ByteSize();
140 EndianConvertReverse(headerSize);
141
142 MessageBuffer packet(sizeof(headerSize) + header.GetCachedSize() + response->GetCachedSize());
143 packet.Write(&headerSize, sizeof(headerSize));
144 uint8* ptr = packet.GetWritePointer();
145 packet.WriteCompleted(header.GetCachedSize());
146 header.SerializePartialToArray(ptr, header.GetCachedSize());
147 ptr = packet.GetWritePointer();
148 packet.WriteCompleted(response->GetCachedSize());
149 response->SerializeToArray(ptr, response->GetCachedSize());
150
151 AsyncWrite(&packet);
152}
153
155{
156 bgs::protocol::Header header;
157 header.set_token(token);
158 header.set_status(status);
159 header.set_service_id(0xFE);
160 header.set_allocated_ciid(&_clientInstanceId);
161
162 auto ciidGuard = Trinity::make_unique_ptr_with_deleter<&bgs::protocol::Header::release_ciid>(&header);
163
164 uint16 headerSize = header.ByteSize();
165 EndianConvertReverse(headerSize);
166
167 MessageBuffer packet(sizeof(headerSize) + header.GetCachedSize());
168 packet.Write(&headerSize, sizeof(headerSize));
169 uint8* ptr = packet.GetWritePointer();
170 packet.WriteCompleted(header.GetCachedSize());
171 header.SerializeToArray(ptr, header.GetCachedSize());
172
173 AsyncWrite(&packet);
174}
175
176void Battlenet::Session::SendRequest(uint32 serviceHash, uint32 methodId, pb::Message const* request, std::function<void(MessageBuffer)> callback)
177{
178 _responseCallbacks[_requestToken] = std::move(callback);
179 SendRequest(serviceHash, methodId, request);
180}
181
182void Battlenet::Session::SendRequest(uint32 serviceHash, uint32 methodId, pb::Message const* request)
183{
184 bgs::protocol::Header header;
185 header.set_service_id(0);
186 header.set_service_hash(serviceHash);
187 header.set_method_id(methodId);
188 header.set_size(request->ByteSize());
189 header.set_token(_requestToken++);
190 header.set_allocated_ciid(&_clientInstanceId);
191
192 auto ciidGuard = Trinity::make_unique_ptr_with_deleter<&bgs::protocol::Header::release_ciid>(&header);
193
194 uint16 headerSize = header.ByteSize();
195 EndianConvertReverse(headerSize);
196
197 MessageBuffer packet(sizeof(headerSize) + header.GetCachedSize() + request->GetCachedSize());
198 packet.Write(&headerSize, sizeof(headerSize));
199 uint8* ptr = packet.GetWritePointer();
200 packet.WriteCompleted(header.GetCachedSize());
201 header.SerializeToArray(ptr, header.GetCachedSize());
202 ptr = packet.GetWritePointer();
203 packet.WriteCompleted(request->GetCachedSize());
204 request->SerializeToArray(ptr, request->GetCachedSize());
205
206 AsyncWrite(&packet);
207}
208
210{
211 _queryProcessor.AddCallback(std::move(queryCallback));
212}
213
214void Battlenet::Session::OnLogon(std::string_view platform, std::string_view locale, uint32 applicationVersion, Minutes timezoneOffset)
215{
216 _locale = locale;
217 _os = platform;
218 _build = applicationVersion;
219
220 _timezoneOffset = timezoneOffset;
221}
222
223void Battlenet::Session::OnLogonSuccess(std::shared_ptr<AccountInfo> accountInfo, std::string_view ipCountry)
224{
225 _accountInfo = std::move(accountInfo);
226 _ipCountry = ipCountry;
227 _authed = true;
228}
229
230void Battlenet::Session::SetClientInfo(uint32 gameAccountId, ClientBuild::VariantId buildVariant, std::array<uint8, 32> const& clientSecret)
231{
232 _gameAccountInfo = GetGameAccountInfo(gameAccountId);
233 _buildVariant = buildVariant;
234 _clientSecret = clientSecret;
235}
236
238{
239 auto itr = _gameAccountInfo->LastPlayedCharacters.find(subRegion);
240 return itr != _gameAccountInfo->LastPlayedCharacters.end() ? &itr->second : nullptr;
241}
242
243template<bool(Battlenet::Session::*processMethod)()>
245{
246 // We have full read header, now check the data payload
247 if (buffer.GetRemainingSpace() > 0)
248 {
249 // need more data in the payload
250 std::size_t readDataSize = std::min(inputBuffer.GetActiveSize(), buffer.GetRemainingSpace());
251 buffer.Write(inputBuffer.GetReadPointer(), readDataSize);
252 inputBuffer.ReadCompleted(readDataSize);
253 }
254
255 if (buffer.GetRemainingSpace() > 0)
256 {
257 // Couldn't receive the whole data this time.
258 ASSERT(inputBuffer.GetActiveSize() == 0);
260 }
261
262 // just received fresh new payload
263 if (!(session->*processMethod)())
264 {
265 session->CloseSocket();
267 }
268
269 return { }; // go to next state
270}
271
273{
274 MessageBuffer& packet = _socket->GetReadBuffer();
275 while (packet.GetActiveSize() > 0)
276 {
277 switch (_packetReadState)
278 {
279 case PacketReadState::HeaderLength:
280 if (Optional<Trinity::Net::SocketReadCallbackResult> partialResult = PartialProcessPacket<&Session::ReadHeaderLengthHandler>(this, packet, _packetBuffer))
281 return *partialResult;
282 [[fallthrough]];
283 case PacketReadState::Header:
284 if (Optional<Trinity::Net::SocketReadCallbackResult> partialResult = PartialProcessPacket<&Session::ReadHeaderHandler>(this, packet, _packetBuffer))
285 return *partialResult;
286 [[fallthrough]];
287 case PacketReadState::Data:
288 if (Optional<Trinity::Net::SocketReadCallbackResult> partialResult = PartialProcessPacket<&Session::ReadDataHandler>(this, packet, _packetBuffer))
289 return *partialResult;
290 break;
291 }
292 }
293
295}
296
298{
299 uint16 len = *reinterpret_cast<uint16*>(_packetBuffer.GetReadPointer());
301
303
304 _packetReadState = PacketReadState::Header;
305 _packetBuffer.Resize(_packetBuffer.GetBufferSize() + len);
306 return true;
307}
308
310{
311 bgs::protocol::Header header;
312 if (!header.ParseFromArray(_packetBuffer.GetReadPointer(), _packetBuffer.GetActiveSize()))
313 return false;
314
315 _packetBuffer.ReadCompleted(_packetBuffer.GetActiveSize());
316
317 _packetReadState = PacketReadState::Data;
318 _packetBuffer.Resize(_packetBuffer.GetBufferSize() + header.size());
319 return true;
320}
321
323{
324 bgs::protocol::Header header;
325 bool parseSuccess = header.ParseFromArray(_packetBuffer.GetBasePointer() + PacketHeaderLengthSize, _packetBuffer.GetReadPointer() - _packetBuffer.GetBasePointer() - PacketHeaderLengthSize);
326 ASSERT(parseSuccess);
327
328 if (header.service_id() != 0xFE)
329 {
330 sServiceDispatcher.Dispatch(this, header.service_hash(), header.token(), header.method_id(), std::move(_packetBuffer));
331 }
332 else
333 {
334 if (auto responseCallback = _responseCallbacks.extract(header.token()))
335 responseCallback.mapped()(std::move(_packetBuffer));
336 else
337 _packetBuffer.Reset();
338 }
339
340 _packetReadState = PacketReadState::HeaderLength;
341 _packetBuffer.Resize(PacketHeaderLengthSize);
342 return true;
343}
344
346{
347 std::ostringstream stream;
348 stream << '[' << _socket->GetRemoteIpAddress() << ':' << _socket->GetRemotePort();
349 if (_accountInfo && !_accountInfo->Login.empty())
350 stream << ", Account: " << _accountInfo->Login;
351
352 if (_gameAccountInfo)
353 stream << ", Game account: " << _gameAccountInfo->Name;
354
355 stream << ']';
356
357 return std::move(stream).str();
358}
void EndianConvertReverse(T &)
AccountTypes
Definition Common.h:42
std::shared_ptr< PreparedResultSet > PreparedQueryResult
uint8_t uint8
Definition Define.h:156
uint16_t uint16
Definition Define.h:155
uint32_t uint32
Definition Define.h:154
std::chrono::system_clock::time_point SystemTimePoint
Definition Duration.h:41
std::chrono::minutes Minutes
Minutes shorthand typedef.
Definition Duration.h:32
#define ASSERT
Definition Errors.h:72
#define TC_LOG_TRACE(filterType__, message__,...)
Definition Log.h:170
std::optional< T > Optional
Optional helper class to wrap optional values within.
Definition Optional.h:25
#define sServiceDispatcher
#define sSessionMgr
constexpr std::size_t PacketHeaderLengthSize
Definition Session.cpp:71
static Optional< Trinity::Net::SocketReadCallbackResult > PartialProcessPacket(Battlenet::Session *session, MessageBuffer &inputBuffer, MessageBuffer &buffer)
Definition Session.cpp:244
std::string GetClientInfo() const
Definition Session.cpp:345
void SendResponse(uint32 token, pb::Message const *response)
Definition Session.cpp:129
void OnLogon(std::string_view platform, std::string_view locale, uint32 applicationVersion, Minutes timezoneOffset)
Definition Session.cpp:214
Trinity::Net::SocketReadCallbackResult ReadHandler()
Definition Session.cpp:272
LastPlayedCharacterInfo const * GetLastPlayedCharacter(std::string_view subRegion) const
Definition Session.cpp:237
void SetClientInfo(uint32 gameAccountId, ClientBuild::VariantId buildVariant, std::array< uint8, 32 > const &clientSecret)
Definition Session.cpp:230
static std::shared_ptr< Socket > CreateSocket(Trinity::Net::IoContextTcpSocket &&socket)
Definition Session.cpp:83
Session(Trinity::Net::IoContextTcpSocket &&socket)
Definition Session.cpp:73
bool ReadHeaderLengthHandler()
Definition Session.cpp:297
void AsyncWrite(MessageBuffer *packet)
Definition Session.cpp:121
void CloseSocket()
Definition Session.h:91
void QueueQuery(QueryCallback &&queryCallback)
Definition Session.cpp:209
void SendRequest(uint32 serviceHash, uint32 methodId, pb::Message const *request, std::function< void(MessageBuffer)> callback)
Definition Session.cpp:176
void OnLogonSuccess(std::shared_ptr< AccountInfo > accountInfo, std::string_view ipCountry)
Definition Session.cpp:223
GameAccountInfo const * GetGameAccountInfo() const
Definition Session.h:104
bool ReadHeaderHandler()
Definition Session.cpp:309
static boost::asio::ssl::context & instance()
Class used to access individual fields of database query result.
Definition Field.h:94
uint64 GetUInt64() const noexcept
Definition Field.cpp:71
bool GetBool() const noexcept
Definition Field.h:102
uint32 GetUInt32() const noexcept
Definition Field.cpp:57
std::string GetString() const noexcept
Definition Field.cpp:113
size_type GetRemainingSpace() const
void ReadCompleted(size_type bytes)
void WriteCompleted(size_type bytes)
uint8 * GetReadPointer()
size_type GetActiveSize() const
uint8 * GetWritePointer()
void Write(void const *data, std::size_t size)
auto MapGetValuePtr(M &map, typename M::key_type const &key)
Definition MapUtils.h:37
SocketReadCallbackResult
Definition Socket.h:49
boost::asio::basic_stream_socket< boost::asio::ip::tcp, Asio::IoContextExecutor > IoContextTcpSocket
Definition Socket.h:40
STL namespace.
void LoadResult(PreparedQueryResult result)
Definition Session.cpp:31
std::string Login
Definition Session.h:70
std::string LockCountry
Definition Session.h:72
std::string LastIP
Definition Session.h:73
std::unordered_map< uint32, GameAccountInfo > GameAccounts
Definition Session.h:78
void LoadResult(Field const *fields)
Definition Session.cpp:53
static std::shared_ptr< SocketConnectionInitializer > & SetupChain(std::span< std::shared_ptr< SocketConnectionInitializer > > initializers)