use crate::{
model::VoiceUpdate,
node::{IncomingEvents, Node, NodeConfig, NodeError, Resume},
player::{Player, PlayerManager},
};
use dashmap::DashMap;
use std::{
error::Error,
fmt::{Display, Formatter, Result as FmtResult},
net::SocketAddr,
sync::Arc,
};
use twilight_model::{
gateway::{event::Event, payload::incoming::VoiceServerUpdate},
id::{
marker::{GuildMarker, UserMarker},
Id,
},
};
#[derive(Debug)]
pub struct ClientError {
kind: ClientErrorType,
source: Option<Box<dyn Error + Send + Sync>>,
}
impl ClientError {
pub const fn kind(&self) -> &ClientErrorType {
&self.kind
}
pub fn into_source(self) -> Option<Box<dyn Error + Send + Sync>> {
self.source
}
#[must_use = "consuming the error into its parts has no effect if left unused"]
pub fn into_parts(self) -> (ClientErrorType, Option<Box<dyn Error + Send + Sync>>) {
(self.kind, self.source)
}
}
impl Display for ClientError {
fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult {
match &self.kind {
ClientErrorType::NodesUnconfigured => f.write_str("no node has been configured"),
ClientErrorType::SendingVoiceUpdate => {
f.write_str("couldn't send voice update to node")
}
}
}
}
impl Error for ClientError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
self.source
.as_ref()
.map(|source| &**source as &(dyn Error + 'static))
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum ClientErrorType {
NodesUnconfigured,
SendingVoiceUpdate,
}
#[derive(Debug)]
pub struct Lavalink {
nodes: DashMap<SocketAddr, Arc<Node>>,
players: PlayerManager,
resume: Option<Resume>,
shard_count: u64,
user_id: Id<UserMarker>,
server_updates: DashMap<Id<GuildMarker>, VoiceServerUpdate>,
sessions: DashMap<Id<GuildMarker>, Box<str>>,
}
impl Lavalink {
pub fn new(user_id: Id<UserMarker>, shard_count: u64) -> Self {
Self::_new_with_resume(user_id, shard_count, None)
}
pub fn new_with_resume(
user_id: Id<UserMarker>,
shard_count: u64,
resume: impl Into<Option<Resume>>,
) -> Self {
Self::_new_with_resume(user_id, shard_count, resume.into())
}
fn _new_with_resume(user_id: Id<UserMarker>, shard_count: u64, resume: Option<Resume>) -> Self {
Self {
nodes: DashMap::new(),
players: PlayerManager::new(),
resume,
shard_count,
user_id,
server_updates: DashMap::new(),
sessions: DashMap::new(),
}
}
pub async fn process(&self, event: &Event) -> Result<(), ClientError> {
tracing::trace!("processing event: {event:?}");
let guild_id = match event {
Event::Ready(e) => {
let shard_id = e.shard.map_or(0, |[id, _]| id);
self.clear_shard_states(shard_id);
return Ok(());
}
Event::VoiceServerUpdate(e) => {
self.server_updates.insert(e.guild_id, e.clone());
e.guild_id
}
Event::VoiceStateUpdate(e) => {
if e.user_id != self.user_id {
tracing::trace!("got voice state update from another user");
return Ok(());
}
if let Some(guild_id) = e.guild_id {
if let Some(player) = self.players.get(&guild_id) {
player.set_channel_id(e.channel_id);
}
if e.channel_id.is_none() {
self.sessions.remove(&guild_id);
} else {
self.sessions
.insert(guild_id, e.session_id.clone().into_boxed_str());
}
guild_id
} else {
tracing::trace!("event has no guild ID: {e:?}");
return Ok(());
}
}
_ => return Ok(()),
};
tracing::debug!("got voice server/state update for {guild_id:?}: {event:?}");
let update = {
let server = self.server_updates.get(&guild_id);
let session = self.sessions.get(&guild_id);
match (server, session) {
(Some(server), Some(session)) => {
let server = server.value();
let session = session.value();
tracing::debug!(
"got both halves for {guild_id}: {server:?}; Session ID: {session:?}",
);
VoiceUpdate::new(guild_id, session.as_ref(), server.clone())
}
(Some(server), None) => {
tracing::debug!(
"guild {guild_id} is now waiting for other half; got: {:?}",
server.value()
);
return Ok(());
}
(None, Some(session)) => {
tracing::debug!(
"guild {guild_id} is now waiting for other half; got session ID: {:?}",
session.value()
);
return Ok(());
}
_ => return Ok(()),
}
};
tracing::debug!("getting player for guild {guild_id}");
let player = self.player(guild_id).await?;
tracing::debug!("sending voice update for guild {guild_id}: {update:?}");
player.send(update).map_err(|source| ClientError {
kind: ClientErrorType::SendingVoiceUpdate,
source: Some(Box::new(source)),
})?;
tracing::debug!("sent voice update for guild {guild_id}");
Ok(())
}
pub async fn add(
&self,
address: SocketAddr,
authorization: impl Into<String>,
) -> Result<(Arc<Node>, IncomingEvents), NodeError> {
let config = NodeConfig {
address,
authorization: authorization.into(),
resume: self.resume.clone(),
user_id: self.user_id,
};
let (node, rx) = Node::connect(config, self.players.clone()).await?;
let node = Arc::new(node);
self.nodes.insert(address, Arc::clone(&node));
Ok((node, rx))
}
pub fn remove(&self, address: SocketAddr) -> Option<(SocketAddr, Arc<Node>)> {
self.nodes.remove(&address)
}
pub fn disconnect(&self, address: SocketAddr) -> bool {
self.nodes.remove(&address).is_some()
}
pub async fn best(&self) -> Result<Arc<Node>, ClientError> {
let mut lowest = i32::MAX;
let mut best = None;
for node in self.nodes.iter() {
if node.sender().is_closed() {
continue;
}
let penalty = node.value().penalty().await;
if penalty < lowest {
lowest = penalty;
best.replace(node.clone());
}
}
best.ok_or(ClientError {
kind: ClientErrorType::NodesUnconfigured,
source: None,
})
}
pub const fn players(&self) -> &PlayerManager {
&self.players
}
pub async fn player(&self, guild_id: Id<GuildMarker>) -> Result<Arc<Player>, ClientError> {
if let Some(player) = self.players().get(&guild_id) {
return Ok(player);
}
let node = self.best().await?;
Ok(self.players().get_or_insert(guild_id, node))
}
fn clear_shard_states(&self, shard_id: u64) {
let shard_count = self.shard_count;
self.server_updates
.retain(|k, _| (k.get() >> 22) % shard_count != shard_id);
self.sessions
.retain(|k, _| (k.get() >> 22) % shard_count != shard_id);
}
}
#[cfg(test)]
mod tests {
use super::{ClientError, ClientErrorType, Lavalink};
use static_assertions::assert_impl_all;
use std::{error::Error, fmt::Debug};
assert_impl_all!(ClientErrorType: Debug, Send, Sync);
assert_impl_all!(ClientError: Error, Send, Sync);
assert_impl_all!(Lavalink: Debug, Send, Sync);
}