handle reconnect messages

This commit is contained in:
ouwou
2020-11-15 02:21:14 -05:00
parent db0833fe86
commit 266fb15c91
7 changed files with 94 additions and 10 deletions

View File

@@ -25,6 +25,7 @@ void DiscordClient::Stop() {
if (!m_client_connected) return;
inflateEnd(&m_zstream);
m_compressed_buf.clear();
m_heartbeat_waiter.kill();
if (m_heartbeat_thread.joinable()) m_heartbeat_thread.join();
@@ -35,6 +36,8 @@ void DiscordClient::Stop() {
m_guild_to_users.clear();
m_websocket.Stop();
m_signal_disconnected.emit(false);
}
bool DiscordClient::IsStarted() const {
@@ -480,15 +483,14 @@ void DiscordClient::HandleGatewayMessage(std::string str) {
try {
switch (m.Opcode) {
case GatewayOp::Hello: {
HelloMessageData d = m.Data;
m_heartbeat_msec = d.HeartbeatInterval;
assert(!m_heartbeat_thread.joinable()); // handle reconnects later
m_heartbeat_thread = std::thread(std::bind(&DiscordClient::HeartbeatThread, this));
SendIdentify();
HandleGatewayHello(m);
} break;
case GatewayOp::HeartbeatAck: {
m_heartbeat_acked = true;
} break;
case GatewayOp::Reconnect: {
HandleGatewayReconnect(m);
} break;
case GatewayOp::Event: {
auto iter = m_event_map.find(m.Type);
if (iter == m_event_map.end()) {
@@ -549,6 +551,18 @@ void DiscordClient::HandleGatewayMessage(std::string str) {
}
}
void DiscordClient::HandleGatewayHello(const GatewayMessage &msg) {
HelloMessageData d = msg.Data;
m_heartbeat_msec = d.HeartbeatInterval;
m_heartbeat_thread = std::thread(std::bind(&DiscordClient::HeartbeatThread, this));
m_signal_connected.emit(); // socket is connected before this but emitting here should b fine
if (m_wants_resume) {
m_wants_resume = false;
SendResume();
} else
SendIdentify();
}
void DiscordClient::ProcessNewGuild(Guild &guild) {
if (guild.IsUnavailable) {
printf("guild (%lld) unavailable\n", static_cast<uint64_t>(guild.ID));
@@ -584,9 +598,10 @@ void DiscordClient::HandleGatewayReady(const GatewayMessage &msg) {
m_store.SetUser(recipient.ID, recipient);
}
m_signal_gateway_ready.emit();
m_session_id = data.SessionID;
m_user_data = data.User;
m_user_settings = data.UserSettings;
m_signal_gateway_ready.emit();
}
void DiscordClient::HandleGatewayMessageCreate(const GatewayMessage &msg) {
@@ -665,6 +680,25 @@ void DiscordClient::HandleGatewayGuildUpdate(const GatewayMessage &msg) {
m_signal_guild_update.emit(id);
}
void DiscordClient::HandleGatewayReconnect(const GatewayMessage &msg) {
m_signal_disconnected.emit(true);
inflateEnd(&m_zstream);
m_compressed_buf.clear();
m_heartbeat_waiter.kill();
if (m_heartbeat_thread.joinable()) m_heartbeat_thread.join();
m_websocket.Stop(1002); // 1000 (kNormalClosureCode) and 1001 will invalidate the session id
std::memset(&m_zstream, 0, sizeof(m_zstream));
inflateInit2(&m_zstream, MAX_WBITS + 32);
m_heartbeat_acked = true;
m_wants_resume = true;
m_websocket.StartConnection(DiscordGateway);
m_websocket.SetMessageCallback(std::bind(&DiscordClient::HandleGatewayMessageRaw, this, std::placeholders::_1));
}
void DiscordClient::HandleGatewayMessageUpdate(const GatewayMessage &msg) {
Snowflake id = msg.Data.at("id");
@@ -774,6 +808,14 @@ void DiscordClient::SendIdentify() {
m_websocket.Send(msg);
}
void DiscordClient::SendResume() {
ResumeMessage msg;
msg.Sequence = m_last_sequence;
msg.SessionID = m_session_id;
msg.Token = m_token;
m_websocket.Send(msg);
}
bool DiscordClient::CheckCode(const cpr::Response &r) {
if (r.status_code >= 300 || r.error) {
fprintf(stderr, "api request to %s failed with status code %d\n", r.url.c_str(), r.status_code);
@@ -843,3 +885,11 @@ DiscordClient::type_signal_channel_create DiscordClient::signal_channel_create()
DiscordClient::type_signal_guild_update DiscordClient::signal_guild_update() {
return m_signal_guild_update;
}
DiscordClient::type_signal_disconnected DiscordClient::signal_disconnected() {
return m_signal_disconnected;
}
DiscordClient::type_signal_connected DiscordClient::signal_connected() {
return m_signal_connected;
}

View File

@@ -118,6 +118,7 @@ private:
void HandleGatewayMessageRaw(std::string str);
void HandleGatewayMessage(std::string str);
void HandleGatewayHello(const GatewayMessage &msg);
void HandleGatewayReady(const GatewayMessage &msg);
void HandleGatewayMessageCreate(const GatewayMessage &msg);
void HandleGatewayMessageDelete(const GatewayMessage &msg);
@@ -132,8 +133,10 @@ private:
void HandleGatewayChannelUpdate(const GatewayMessage &msg);
void HandleGatewayChannelCreate(const GatewayMessage &msg);
void HandleGatewayGuildUpdate(const GatewayMessage &msg);
void HandleGatewayReconnect(const GatewayMessage &msg);
void HeartbeatThread();
void SendIdentify();
void SendResume();
bool CheckCode(const cpr::Response &r);
@@ -165,6 +168,9 @@ private:
HeartbeatWaiter m_heartbeat_waiter;
std::atomic<bool> m_heartbeat_acked = true;
bool m_wants_resume = false;
std::string m_session_id;
mutable std::mutex m_msg_mutex;
Glib::Dispatcher m_msg_dispatch;
std::queue<std::string> m_msg_queue;
@@ -183,6 +189,8 @@ public:
typedef sigc::signal<void, Snowflake> type_signal_channel_update;
typedef sigc::signal<void, Snowflake> type_signal_channel_create;
typedef sigc::signal<void, Snowflake> type_signal_guild_update;
typedef sigc::signal<void, bool> type_signal_disconnected; // bool true if reconnecting
typedef sigc::signal<void> type_signal_connected;
type_signal_gateway_ready signal_gateway_ready();
type_signal_message_create signal_message_create();
@@ -195,6 +203,8 @@ public:
type_signal_channel_update signal_channel_update();
type_signal_channel_create signal_channel_create();
type_signal_guild_update signal_guild_update();
type_signal_disconnected signal_disconnected();
type_signal_connected signal_connected();
protected:
type_signal_gateway_ready m_signal_gateway_ready;
@@ -208,4 +218,6 @@ protected:
type_signal_channel_update m_signal_channel_update;
type_signal_channel_create m_signal_channel_create;
type_signal_guild_update m_signal_guild_update;
type_signal_disconnected m_signal_disconnected;
type_signal_connected m_signal_connected;
};

View File

@@ -171,3 +171,11 @@ void to_json(nlohmann::json &j, const CreateDMObject &m) {
conv.push_back(std::to_string(id));
j["recipients"] = conv;
}
void to_json(nlohmann::json &j, const ResumeMessage &m) {
j["op"] = GatewayOp::Resume;
j["d"] = nlohmann::json::object();
j["d"]["token"] = m.Token;
j["d"]["session_id"] = m.SessionID;
j["d"]["seq"] = m.Sequence;
}

View File

@@ -24,6 +24,8 @@ enum class GatewayOp : int {
Heartbeat = 1,
Identify = 2,
UpdateStatus = 3,
Resume = 6,
Reconnect = 7,
Hello = 10,
HeartbeatAck = 11,
LazyLoadRequest = 14,
@@ -241,3 +243,11 @@ struct CreateDMObject {
friend void to_json(nlohmann::json &j, const CreateDMObject &m);
};
struct ResumeMessage : GatewayMessage {
std::string Token;
std::string SessionID;
int Sequence;
friend void to_json(nlohmann::json &j, const ResumeMessage &m);
};

View File

@@ -178,8 +178,11 @@ const Store::roles_type &Store::GetRoles() const {
void Store::ClearAll() {
m_channels.clear();
m_emojis.clear();
m_guilds.clear();
m_members.clear();
m_messages.clear();
m_permissions.clear();
m_roles.clear();
m_users.clear();
}

View File

@@ -14,6 +14,10 @@ void Websocket::Stop() {
m_websocket.stop();
}
void Websocket::Stop(uint16_t code) {
m_websocket.stop(code);
}
bool Websocket::IsOpen() const {
auto state = m_websocket.getReadyState();
return state == ix::ReadyState::Open;
@@ -35,10 +39,6 @@ void Websocket::Send(const nlohmann::json &j) {
void Websocket::OnMessage(const ix::WebSocketMessagePtr &msg) {
switch (msg->type) {
case ix::WebSocketMessageType::Message: {
//if (msg->str.size() > 1000)
// printf("%s\n", msg->str.substr(0, 1000).c_str());
//else
// printf("%s\n", msg->str.c_str());
if (m_callback)
m_callback(msg->str);
} break;

View File

@@ -15,6 +15,7 @@ public:
void Send(const std::string &str);
void Send(const nlohmann::json &j);
void Stop();
void Stop(uint16_t code);
bool IsOpen() const;
private: