use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use futures::channel::mpsc::{UnboundedReceiver as Receiver, UnboundedSender as Sender};
use futures::StreamExt;
use tokio::sync::{Mutex, RwLock};
use tokio::time::{sleep, timeout, Duration, Instant};
use tracing::{debug, info, instrument, warn};
use typemap_rev::TypeMap;
use super::{
ShardClientMessage,
ShardId,
ShardManagerMessage,
ShardMessenger,
ShardQueuerMessage,
ShardRunner,
ShardRunnerInfo,
ShardRunnerOptions,
};
#[cfg(feature = "voice")]
use crate::client::bridge::voice::VoiceGatewayManager;
use crate::client::{EventHandler, RawEventHandler};
#[cfg(feature = "framework")]
use crate::framework::Framework;
use crate::gateway::{ConnectionStage, InterMessage, Shard};
use crate::internal::prelude::*;
use crate::internal::tokio::spawn_named;
use crate::model::gateway::GatewayIntents;
use crate::CacheAndHttp;
const WAIT_BETWEEN_BOOTS_IN_SECONDS: u64 = 5;
pub struct ShardQueuer {
pub data: Arc<RwLock<TypeMap>>,
pub event_handler: Option<Arc<dyn EventHandler>>,
pub raw_event_handler: Option<Arc<dyn RawEventHandler>>,
#[cfg(feature = "framework")]
pub framework: Arc<dyn Framework + Send + Sync>,
pub last_start: Option<Instant>,
pub manager_tx: Sender<ShardManagerMessage>,
pub queue: VecDeque<(u64, u64)>,
pub runners: Arc<Mutex<HashMap<ShardId, ShardRunnerInfo>>>,
pub rx: Receiver<ShardQueuerMessage>,
#[cfg(feature = "voice")]
pub voice_manager: Option<Arc<dyn VoiceGatewayManager + Send + Sync + 'static>>,
pub ws_url: Arc<Mutex<String>>,
pub cache_and_http: Arc<CacheAndHttp>,
pub intents: GatewayIntents,
}
impl ShardQueuer {
#[instrument(skip(self))]
pub async fn run(&mut self) {
const TIMEOUT: Duration = Duration::from_secs(WAIT_BETWEEN_BOOTS_IN_SECONDS);
loop {
match timeout(TIMEOUT, self.rx.next()).await {
Ok(Some(ShardQueuerMessage::Shutdown)) => {
debug!("[Shard Queuer] Received to shutdown.");
self.shutdown_runners().await;
break;
},
Ok(Some(ShardQueuerMessage::ShutdownShard(shard, code))) => {
debug!("[Shard Queuer] Received to shutdown shard {} with {}.", shard.0, code);
self.shutdown(shard, code).await;
},
Ok(Some(ShardQueuerMessage::Start(id, total))) => {
debug!("[Shard Queuer] Received to start shard {} of {}.", id.0, total.0);
self.checked_start(id.0, total.0).await;
},
Ok(None) => break,
Err(_) => {
if let Some((id, total)) = self.queue.pop_front() {
self.checked_start(id, total).await;
}
},
}
}
}
#[instrument(skip(self))]
async fn check_last_start(&mut self) {
let instant = match self.last_start {
Some(instant) => instant,
None => return,
};
let duration = Duration::from_secs(WAIT_BETWEEN_BOOTS_IN_SECONDS);
let elapsed = instant.elapsed();
if elapsed >= duration {
return;
}
let to_sleep = duration - elapsed;
sleep(to_sleep).await;
}
#[instrument(skip(self))]
async fn checked_start(&mut self, id: u64, total: u64) {
debug!("[Shard Queuer] Checked start for shard {} out of {}", id, total);
self.check_last_start().await;
if let Err(why) = self.start(id, total).await {
warn!("[Shard Queuer] Err starting shard {}: {:?}", id, why);
info!("[Shard Queuer] Re-queueing start of shard {}", id);
self.queue.push_back((id, total));
}
self.last_start = Some(Instant::now());
}
#[instrument(skip(self))]
async fn start(&mut self, shard_id: u64, shard_total: u64) -> Result<()> {
let shard_info = [shard_id, shard_total];
let mut shard = Shard::new(
Arc::clone(&self.ws_url),
&self.cache_and_http.http.token,
shard_info,
self.intents,
)
.await?;
shard.set_http(Arc::clone(&self.cache_and_http.http));
let mut runner = ShardRunner::new(ShardRunnerOptions {
data: Arc::clone(&self.data),
event_handler: self.event_handler.as_ref().map(Arc::clone),
raw_event_handler: self.raw_event_handler.as_ref().map(Arc::clone),
#[cfg(feature = "framework")]
framework: Arc::clone(&self.framework),
manager_tx: self.manager_tx.clone(),
#[cfg(feature = "voice")]
voice_manager: self.voice_manager.clone(),
shard,
cache_and_http: Arc::clone(&self.cache_and_http),
});
let runner_info = ShardRunnerInfo {
latency: None,
runner_tx: ShardMessenger::new(runner.runner_tx()),
stage: ConnectionStage::Disconnected,
};
spawn_named("shard_queuer::stop", async move {
drop(runner.run().await);
debug!("[ShardRunner {:?}] Stopping", runner.shard.shard_info());
});
self.runners.lock().await.insert(ShardId(shard_id), runner_info);
Ok(())
}
#[instrument(skip(self))]
async fn shutdown_runners(&mut self) {
let keys = {
let runners = self.runners.lock().await;
if runners.is_empty() {
return;
}
runners.keys().copied().collect::<Vec<_>>()
};
info!("Shutting down all shards");
for shard_id in keys {
self.shutdown(shard_id, 1000).await;
}
}
#[instrument(skip(self))]
pub async fn shutdown(&mut self, shard_id: ShardId, code: u16) {
info!("Shutting down shard {}", shard_id);
if let Some(runner) = self.runners.lock().await.get(&shard_id) {
let shutdown = ShardManagerMessage::Shutdown(shard_id, code);
let client_msg = ShardClientMessage::Manager(shutdown);
let msg = InterMessage::Client(Box::new(client_msg));
if let Err(why) = runner.runner_tx.tx.unbounded_send(msg) {
warn!(
"Failed to cleanly shutdown shard {} when sending message to shard runner: {:?}",
shard_id,
why,
);
}
}
}
}