use crate::constants::{self, close_codes};
use crate::internal::prelude::*;
use crate::model::{
event::{Event, GatewayEvent},
gateway::Activity,
id::GuildId,
user::OnlineStatus
};
use tokio::sync::Mutex;
use crate::client::bridge::gateway::{GatewayIntents, ChunkGuildFilter};
use std::{
sync::Arc,
time::{Duration as StdDuration, Instant}
};
use super::{
ConnectionStage,
CurrentPresence,
ShardAction,
GatewayError,
ReconnectType,
WsStream,
WebSocketGatewayClientExt,
};
use async_tungstenite::tungstenite::{
error::Error as TungsteniteError,
protocol::frame::CloseFrame,
};
use url::Url;
use tracing::{error, debug, info, trace, warn, instrument};
#[cfg(all(feature = "rustls_backend", not(feature = "native_tls_backend")))]
use crate::internal::ws_impl::create_rustls_client;
#[cfg(feature = "native_tls_backend")]
use crate::internal::ws_impl::create_native_tls_client;
pub struct Shard {
pub client: WsStream,
current_presence: CurrentPresence,
heartbeat_instants: (Option<Instant>, Option<Instant>),
heartbeat_interval: Option<u64>,
last_heartbeat_acknowledged: bool,
seq: u64,
session_id: Option<String>,
shard_info: [u64; 2],
guild_subscriptions: bool,
shutdown: bool,
stage: ConnectionStage,
pub started: Instant,
pub token: String,
ws_url: Arc<Mutex<String>>,
pub intents: Option<GatewayIntents>,
}
impl Shard {
pub async fn new(
ws_url: Arc<Mutex<String>>,
token: &str,
shard_info: [u64; 2],
guild_subscriptions: bool,
intents: Option<GatewayIntents>,
) -> Result<Shard> {
let url = ws_url.lock().await.clone();
let client = connect(&url).await?;
let current_presence = (None, OnlineStatus::Online);
let heartbeat_instants = (None, None);
let heartbeat_interval = None;
let last_heartbeat_acknowledged = true;
let seq = 0;
let stage = ConnectionStage::Handshake;
let session_id = None;
Ok(Shard {
shutdown: false,
client,
current_presence,
heartbeat_instants,
heartbeat_interval,
last_heartbeat_acknowledged,
seq,
stage,
started: Instant::now(),
token: token.to_string(),
session_id,
shard_info,
guild_subscriptions,
ws_url,
intents,
})
}
#[inline]
pub fn current_presence(&self) -> &CurrentPresence {
&self.current_presence
}
#[inline]
pub fn is_shutdown(&self) -> bool {
self.shutdown
}
#[inline]
pub fn heartbeat_instants(&self) -> &(Option<Instant>, Option<Instant>) {
&self.heartbeat_instants
}
#[inline]
pub fn last_heartbeat_sent(&self) -> Option<&Instant> {
self.heartbeat_instants.0.as_ref()
}
#[inline]
pub fn last_heartbeat_ack(&self) -> Option<&Instant> {
self.heartbeat_instants.1.as_ref()
}
#[instrument(skip(self))]
pub async fn heartbeat(&mut self) -> Result<()> {
match self.client.send_heartbeat(&self.shard_info, Some(self.seq)).await {
Ok(()) => {
self.heartbeat_instants.0 = Some(Instant::now());
self.last_heartbeat_acknowledged = false;
Ok(())
},
Err(why) => {
match why {
Error::Tungstenite(TungsteniteError::Io(err)) => if err.raw_os_error() != Some(32) {
debug!("[Shard {:?}] Err heartbeating: {:?}",
self.shard_info,
err);
},
other => {
warn!("[Shard {:?}] Other err w/ keepalive: {:?}",
self.shard_info,
other);
},
}
Err(Error::Gateway(GatewayError::HeartbeatFailed))
}
}
}
#[inline]
pub fn heartbeat_interval(&self) -> Option<&u64> {
self.heartbeat_interval.as_ref()
}
#[inline]
pub fn last_heartbeat_acknowledged(&self) -> bool {
self.last_heartbeat_acknowledged
}
#[inline]
pub fn seq(&self) -> u64 {
self.seq
}
#[inline]
pub fn session_id(&self) -> Option<&String> {
self.session_id.as_ref()
}
#[inline]
#[instrument(skip(self))]
pub fn set_activity(&mut self, activity: Option<Activity>) {
self.current_presence.0 = activity;
}
#[inline]
#[instrument(skip(self))]
pub fn set_presence(&mut self, status: OnlineStatus, activity: Option<Activity>) {
self.set_activity(activity);
self.set_status(status);
}
#[inline]
#[instrument(skip(self))]
pub fn set_status(&mut self, mut status: OnlineStatus) {
if status == OnlineStatus::Offline {
status = OnlineStatus::Invisible;
}
self.current_presence.1 = status;
}
pub fn shard_info(&self) -> [u64; 2] { self.shard_info }
pub fn stage(&self) -> ConnectionStage {
self.stage
}
#[instrument(skip(self))]
fn handle_gateway_dispatch(&mut self, seq: u64, event: &Event) -> Result<Option<ShardAction>> {
if seq > self.seq + 1 {
warn!("[Shard {:?}] Sequence off; them: {}, us: {}", self.shard_info, seq, self.seq);
}
match event {
Event::Ready(ref ready) => {
debug!("[Shard {:?}] Received Ready", self.shard_info);
self.session_id = Some(ready.ready.session_id.clone());
self.stage = ConnectionStage::Connected;
},
Event::Resumed(_) => {
info!("[Shard {:?}] Resumed", self.shard_info);
self.stage = ConnectionStage::Connected;
self.last_heartbeat_acknowledged = true;
self.heartbeat_instants = (Some(Instant::now()), None);
},
_ => {},
}
self.seq = seq;
Ok(None)
}
#[instrument(skip(self))]
fn handle_heartbeat_event(&mut self, s: u64) -> Result<Option<ShardAction>> {
info!("[Shard {:?}] Received shard heartbeat", self.shard_info);
if s > self.seq + 1 {
info!(
"[Shard {:?}] Received off sequence (them: {}; us: {}); resuming",
self.shard_info,
s,
self.seq
);
if self.stage == ConnectionStage::Handshake {
self.stage = ConnectionStage::Identifying;
return Ok(Some(ShardAction::Identify));
} else {
warn!(
"[Shard {:?}] Heartbeat during non-Handshake; auto-reconnecting",
self.shard_info
);
return Ok(Some(ShardAction::Reconnect(self.reconnection_type())));
}
}
Ok(Some(ShardAction::Heartbeat))
}
#[instrument(skip(self))]
fn handle_gateway_closed(&mut self, data: &Option<CloseFrame<'static>>) -> Result<Option<ShardAction>> {
let num = data.as_ref().map(|d| d.code.into());
let clean = num == Some(1000);
match num {
Some(close_codes::UNKNOWN_OPCODE) => {
warn!("[Shard {:?}] Sent invalid opcode.",
self.shard_info);
},
Some(close_codes::DECODE_ERROR) => {
warn!("[Shard {:?}] Sent invalid message.",
self.shard_info);
},
Some(close_codes::NOT_AUTHENTICATED) => {
warn!("[Shard {:?}] Sent no authentication.",
self.shard_info);
return Err(Error::Gateway(GatewayError::NoAuthentication));
},
Some(close_codes::AUTHENTICATION_FAILED) => {
error!("[Shard {:?}] Sent invalid authentication, please check the token.", self.shard_info);
return Err(Error::Gateway(GatewayError::InvalidAuthentication));
},
Some(close_codes::ALREADY_AUTHENTICATED) => {
warn!("[Shard {:?}] Already authenticated.",
self.shard_info);
},
Some(close_codes::INVALID_SEQUENCE) => {
warn!("[Shard {:?}] Sent invalid seq: {}.",
self.shard_info,
self.seq);
self.seq = 0;
},
Some(close_codes::RATE_LIMITED) => {
warn!("[Shard {:?}] Gateway ratelimited.",
self.shard_info);
},
Some(close_codes::INVALID_SHARD) => {
warn!("[Shard {:?}] Sent invalid shard data.",
self.shard_info);
return Err(Error::Gateway(GatewayError::InvalidShardData));
},
Some(close_codes::SHARDING_REQUIRED) => {
error!("[Shard {:?}] Shard has too many guilds.",
self.shard_info);
return Err(Error::Gateway(GatewayError::OverloadedShard));
},
Some(4006) | Some(close_codes::SESSION_TIMEOUT) => {
info!("[Shard {:?}] Invalid session.", self.shard_info);
self.session_id = None;
},
Some(close_codes::INVALID_GATEWAY_INTENTS) => {
error!("[Shard {:?}] Invalid gateway intents have been provided.", self.shard_info);
return Err(Error::Gateway(GatewayError::InvalidGatewayIntents));
},
Some(close_codes::DISALLOWED_GATEWAY_INTENTS) => {
error!("[Shard {:?}] Disallowed gateway intents have been provided.", self.shard_info);
return Err(Error::Gateway(GatewayError::DisallowedGatewayIntents));
},
Some(other) if !clean => {
warn!(
"[Shard {:?}] Unknown unclean close {}: {:?}",
self.shard_info,
other,
data.as_ref().map(|d| &d.reason),
);
},
_ => {},
}
let resume = num.map(|x| {
x != close_codes::AUTHENTICATION_FAILED &&
self.session_id.is_some()
}).unwrap_or(true);
Ok(Some(if resume {
ShardAction::Reconnect(ReconnectType::Resume)
} else {
ShardAction::Reconnect(ReconnectType::Reidentify)
}))
}
#[instrument(skip(self))]
pub(crate) fn handle_event(&mut self, event: &Result<GatewayEvent>)
-> Result<Option<ShardAction>> {
match *event {
Ok(GatewayEvent::Dispatch(seq, ref event)) => self.handle_gateway_dispatch(seq, event),
Ok(GatewayEvent::Heartbeat(s)) => self.handle_heartbeat_event(s),
Ok(GatewayEvent::HeartbeatAck) => {
self.heartbeat_instants.1 = Some(Instant::now());
self.last_heartbeat_acknowledged = true;
trace!("[Shard {:?}] Received heartbeat ack", self.shard_info);
Ok(None)
},
Ok(GatewayEvent::Hello(interval)) => {
debug!("[Shard {:?}] Received a Hello; interval: {}",
self.shard_info,
interval);
if self.stage == ConnectionStage::Resuming {
return Ok(None);
}
if interval > 0 {
self.heartbeat_interval = Some(interval);
}
Ok(Some(if self.stage == ConnectionStage::Handshake {
ShardAction::Identify
} else {
debug!("[Shard {:?}] Received late Hello; autoreconnecting",
self.shard_info);
ShardAction::Reconnect(self.reconnection_type())
}))
},
Ok(GatewayEvent::InvalidateSession(resumable)) => {
info!(
"[Shard {:?}] Received session invalidation",
self.shard_info,
);
Ok(Some(if resumable {
ShardAction::Reconnect(ReconnectType::Resume)
} else {
ShardAction::Reconnect(ReconnectType::Reidentify)
}))
},
Ok(GatewayEvent::Reconnect) => {
Ok(Some(ShardAction::Reconnect(ReconnectType::Resume)))
},
Err(Error::Gateway(GatewayError::Closed(ref data))) => self.handle_gateway_closed(&data),
Err(Error::Tungstenite(ref why)) => {
warn!("[Shard {:?}] Websocket error: {:?}",
self.shard_info,
why);
info!("[Shard {:?}] Will attempt to auto-reconnect",
self.shard_info);
Ok(Some(ShardAction::Reconnect(self.reconnection_type())))
},
_ => Ok(None),
}
}
#[instrument(skip(self))]
pub async fn check_heartbeat(&mut self) -> bool {
let wait = {
let heartbeat_interval = match self.heartbeat_interval {
Some(heartbeat_interval) => heartbeat_interval,
None => {
return self.started.elapsed() < StdDuration::from_secs(15);
},
};
StdDuration::from_secs(heartbeat_interval / 1000)
};
if let Some(last_sent) = self.heartbeat_instants.0 {
if last_sent.elapsed() <= wait {
return true;
}
}
if !self.last_heartbeat_acknowledged {
debug!(
"[Shard {:?}] Last heartbeat not acknowledged",
self.shard_info,
);
return false;
}
if let Err(why) = self.heartbeat().await {
warn!("[Shard {:?}] Err heartbeating: {:?}", self.shard_info, why);
false
} else {
trace!("[Shard {:?}] Heartbeat", self.shard_info);
true
}
}
#[instrument(skip(self))]
pub fn latency(&self) -> Option<StdDuration> {
if let (Some(sent), Some(received)) = self.heartbeat_instants {
if received > sent {
return Some(received - sent);
}
}
None
}
pub fn should_reconnect(&mut self) -> Option<ReconnectType> {
if self.stage == ConnectionStage::Connecting {
return None;
}
Some(self.reconnection_type())
}
pub fn reconnection_type(&self) -> ReconnectType {
if self.session_id().is_some() {
ReconnectType::Resume
} else {
ReconnectType::Reidentify
}
}
#[instrument(skip(self))]
pub async fn chunk_guild(
&mut self,
guild_id: GuildId,
limit: Option<u16>,
filter: ChunkGuildFilter,
nonce: Option<&str>,
) -> Result<()> {
debug!("[Shard {:?}] Requesting member chunks", self.shard_info);
self.client.send_chunk_guild(
guild_id,
&self.shard_info,
limit,
filter,
nonce,
).await
}
#[instrument(skip(self))]
pub async fn identify(&mut self) -> Result<()> {
self.client.send_identify(&self.shard_info, &self.token, self.guild_subscriptions, self.intents).await?;
self.heartbeat_instants.0 = Some(Instant::now());
self.stage = ConnectionStage::Identifying;
Ok(())
}
#[instrument(skip(self))]
pub async fn initialize(&mut self) -> Result<WsStream> {
debug!("[Shard {:?}] Initializing.", self.shard_info);
self.stage = ConnectionStage::Connecting;
self.started = Instant::now();
let url = &self.ws_url.lock().await.clone();
let client = connect(&url).await?;
self.stage = ConnectionStage::Handshake;
Ok(client)
}
#[instrument(skip(self))]
pub async fn reset(&mut self) {
self.heartbeat_instants = (Some(Instant::now()), None);
self.heartbeat_interval = None;
self.last_heartbeat_acknowledged = true;
self.session_id = None;
self.stage = ConnectionStage::Disconnected;
self.seq = 0;
}
#[instrument(skip(self))]
pub async fn resume(&mut self) -> Result<()> {
debug!("[Shard {:?}] Attempting to resume", self.shard_info);
self.client = self.initialize().await?;
self.stage = ConnectionStage::Resuming;
match self.session_id.as_ref() {
Some(session_id) => {
self.client.send_resume(
&self.shard_info,
session_id,
self.seq,
&self.token,
).await
},
None => Err(Error::Gateway(GatewayError::NoSessionId)),
}
}
#[instrument(skip(self))]
pub async fn reconnect(&mut self) -> Result<()> {
info!("[Shard {:?}] Attempting to reconnect", self.shard_info());
self.reset().await;
self.client = self.initialize().await?;
Ok(())
}
#[instrument(skip(self))]
pub async fn update_presence(&mut self) -> Result<()> {
self.client.send_presence_update(
&self.shard_info,
&self.current_presence,
).await
}
}
#[cfg(all(feature = "rustls_backend", not(feature = "native_tls_backend")))]
async fn connect(base_url: &str) -> Result<WsStream> {
let url = build_gateway_url(base_url)?;
Ok(create_rustls_client(url).await?)
}
#[cfg(feature = "native_tls_backend")]
async fn connect(base_url: &str) -> Result<WsStream> {
let url = build_gateway_url(base_url)?;
Ok(create_native_tls_client(url).await?)
}
fn build_gateway_url(base: &str) -> Result<Url> {
Url::parse(&format!("{}?v={}", base, constants::GATEWAY_VERSION))
.map_err(|why| {
warn!("Error building gateway URL with base `{}`: {:?}", base, why);
Error::Gateway(GatewayError::BuildingUrl)
})
}