refactor: use mpsc instead of locking

This commit is contained in:
2022-08-15 01:14:35 +02:00
parent 1f9eed55c4
commit 4e85085b06
3 changed files with 28 additions and 11 deletions
Generated
+1
View File
@@ -965,6 +965,7 @@ name = "spamcontest"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"chrono", "chrono",
"dashmap",
"env_logger", "env_logger",
"itertools", "itertools",
"log", "log",
+1
View File
@@ -11,6 +11,7 @@ env_logger = "0.9"
chrono = "0.4" chrono = "0.4"
itertools = "0.10" itertools = "0.10"
tokio = {version = "1.20", features = ["rt-multi-thread", "signal"]} tokio = {version = "1.20", features = ["rt-multi-thread", "signal"]}
dashmap = "5.3.4"
[dependencies.serenity] [dependencies.serenity]
version = "0.11" version = "0.11"
+26 -11
View File
@@ -1,3 +1,4 @@
use dashmap::DashMap;
use itertools::Itertools; use itertools::Itertools;
use log::{debug, error, info}; use log::{debug, error, info};
use serenity::client::{Context, EventHandler}; use serenity::client::{Context, EventHandler};
@@ -10,8 +11,8 @@ use std::cmp::Reverse;
use std::collections::HashMap; use std::collections::HashMap;
use std::fmt::Display; use std::fmt::Display;
use std::ops::RangeInclusive; use std::ops::RangeInclusive;
use std::sync::Mutex;
use std::time::Duration; use std::time::Duration;
use tokio::sync::mpsc;
const DEFAULT_CONTEST_DURATION: Duration = Duration::from_secs(60); const DEFAULT_CONTEST_DURATION: Duration = Duration::from_secs(60);
@@ -20,7 +21,7 @@ const ALLOWED_DURATION_RANGE: RangeInclusive<Duration> =
const PIN_ANNOUNCEMENT_THRESHOLD: Duration = Duration::from_secs(5 * 60); const PIN_ANNOUNCEMENT_THRESHOLD: Duration = Duration::from_secs(5 * 60);
type Contests = Mutex<HashMap<ChannelId, Contest>>; type Contests = DashMap<ChannelId, mpsc::Sender<Message>>;
#[derive(Default)] #[derive(Default)]
pub struct Handler { pub struct Handler {
@@ -36,14 +37,14 @@ impl Handler {
#[serenity::async_trait] #[serenity::async_trait]
impl EventHandler for Handler { impl EventHandler for Handler {
async fn message(&self, ctx: Context, msg: Message) { async fn message(&self, ctx: Context, msg: Message) {
if let Some(contest) = self.contests.lock().unwrap().get_mut(&msg.channel_id) { if let Some(contest) = self.contests.get(&msg.channel_id) {
debug!( debug!(
"Counting message {} (from {} in channel {})", "Counting message {} (from {} in channel {})",
msg.id, msg.id,
msg.author.tag(), msg.author.tag(),
msg.channel_id.0 msg.channel_id.0
); );
contest.count(&msg); contest.value().send(msg).await.unwrap();
return; return;
} }
@@ -183,11 +184,25 @@ async fn run_contest(
announcement.pin(&ctx.http).await.ok(); announcement.pin(&ctx.http).await.ok();
} }
contests.lock().unwrap().insert(channel_id, Contest::new()); let mut counts = Contest::new();
tokio::time::sleep(duration).await;
let contest = contests.lock().unwrap().remove(&channel_id).unwrap();
if contest.counts.is_empty() { {
let (tx, mut rx) = mpsc::channel(8);
contests.insert(channel_id, tx);
tokio::select! {
_ = tokio::time::sleep(duration) => {},
_ = async {
while let Some(msg) = rx.recv().await {
counts.count(&msg);
}
} => { unreachable!("mpsc receiver closed unexpectedly") },
}
contests.remove(&channel_id);
}
if counts.counts.is_empty() {
announcement.delete(&ctx.http).await?; announcement.delete(&ctx.http).await?;
} else { } else {
if pin { if pin {
@@ -202,12 +217,12 @@ async fn run_contest(
.colour(Colour::DARK_GREEN) .colour(Colour::DARK_GREEN)
.field( .field(
"Ergebnisse (nach Nachrichten):", "Ergebnisse (nach Nachrichten):",
contest.ranking_by(|c| Reverse(c.messages), |c| c.messages), counts.ranking_by(|c| Reverse(c.messages), |c| c.messages),
false, false,
) )
.field( .field(
"Ergebnisse (nach Zeichen):", "Ergebnisse (nach Zeichen):",
contest.ranking_by(|c| Reverse(c.characters), |c| c.characters), counts.ranking_by(|c| Reverse(c.characters), |c| c.characters),
false, false,
) )
}) })
@@ -215,5 +230,5 @@ async fn run_contest(
.await?; .await?;
} }
Ok(contest) Ok(counts)
} }