diff --git a/README.md b/README.md index eb95fed..2a24308 100644 --- a/README.md +++ b/README.md @@ -39,6 +39,7 @@ dev portal to make multiple people owner. - `/markov`: For a reply from Bingus in the current channel - `/weights`: View the weights for a specific token +- *`/forget`: Clear a word from memory, this gets rid of any instance of the token - *`/dump_chain`: Dump Bingus' "brain", his entire database of known words and relations - *`/load_chain`: Additively load a brain file into Bingus diff --git a/src/brain.rs b/src/brain.rs index f337cef..4999ae1 100644 --- a/src/brain.rs +++ b/src/brain.rs @@ -65,6 +65,12 @@ impl Edges { None } + pub fn forget(&mut self, token: &Token) { + if let Some(w) = self.0.remove(token) { + self.1 -= w as u64; + } + } + pub fn iter_weights(&self) -> impl Iterator { self.0 .iter() @@ -137,6 +143,16 @@ impl Brain { .unwrap_or_default() } + pub fn forget(&mut self, word: &str) { + let tok = Self::normalize_token(word); + + self.0.remove(&tok); + + for edge in self.0.values_mut() { + edge.forget(&tok); + } + } + pub fn merge_from(&mut self, other: Self) { for (k, v) in other.0.into_iter() { if let Some(edges) = self.0.get_mut(&k) { @@ -226,7 +242,9 @@ impl Brain { } pub fn get_weights(&self, tok: &str) -> Option<&Edges> { - self.0.get(&Self::normalize_token(tok)) + self.0 + .get(&Self::normalize_token(tok)) + .filter(|e| !e.0.is_empty()) } fn legacy_token_format(tok: &Token) -> String { @@ -338,6 +356,30 @@ mod tests { } } + #[test] + fn forget_word() { + let mut brain = Brain::default(); + + brain.ingest("hello world"); + brain.ingest("hello evil world"); + + brain.forget("evil"); + + assert!( + !brain.0.contains_key(&Some(String::from("evil"))), + "Edges still exist for evil" + ); + let edges = brain + .0 + .get(&Some(String::from("hello"))) + .expect("No weights for hello"); + assert!( + !edges.0.contains_key(&Some(String::from("evil"))), + "Edges for hello still has evil" + ); + assert_eq!(edges.1, 1); + } + #[test] fn none_on_empty() { let mut brain = Brain::default(); diff --git a/src/cmd/forget.rs b/src/cmd/forget.rs new file mode 100644 index 0000000..9fd1132 --- /dev/null +++ b/src/cmd/forget.rs @@ -0,0 +1,45 @@ +use std::sync::Arc; + +use twilight_interactions::command::{CommandModel, CreateCommand}; +use twilight_model::application::interaction::{Interaction, application_command::CommandData}; + +use crate::{BotContext, cmd::DEFER_INTER_RESP_EPHEMERAL, prelude::*, require_owner}; + +#[derive(CommandModel, CreateCommand)] +#[command( + name = "forget", + desc = "Erase a word from all edges. THIS ACTION IS IRREVERSIBLE!" +)] +pub struct ForgetCommand { + /// The token to forget + token: String, +} + +impl ForgetCommand { + pub async fn handle(inter: Interaction, data: CommandData, ctx: Arc) -> Result { + let client = ctx.http.interaction(ctx.app_id); + + require_owner!(inter, ctx, client); + + let Self { token } = + Self::from_interaction(data.into()).context("Failed to parse command data")?; + + client + .create_response(inter.id, &inter.token, &DEFER_INTER_RESP_EPHEMERAL) + .await + .context("Failed to defer")?; + + { + let mut brain = ctx.brain_handle.write().await; + brain.forget(token.as_str()); + } + + client + .update_response(&inter.token) + .content(Some("Token forgotten")) + .await + .context("Failed to send brain")?; + + Ok(()) + } +} diff --git a/src/cmd/mod.rs b/src/cmd/mod.rs index 02160e0..4b7c0f1 100644 --- a/src/cmd/mod.rs +++ b/src/cmd/mod.rs @@ -1,4 +1,5 @@ mod dump_chain; +mod forget; mod load_chain; mod markov; mod weights; @@ -16,6 +17,7 @@ use twilight_model::http::interaction::{ use crate::{BotContext, prelude::*}; use dump_chain::DumpChainCommand; +use forget::ForgetCommand; use load_chain::LoadChainCommand; use markov::MarkovCommand; use weights::WeightsCommand; @@ -66,6 +68,7 @@ pub async fn register_all_commands(ctx: Arc) -> Result { DumpChainCommand::create_command().into(), LoadChainCommand::create_command().into(), MarkovCommand::create_command().into(), + ForgetCommand::create_command().into(), ]; let client = ctx.http.interaction(ctx.app_id); @@ -88,6 +91,7 @@ pub async fn handle_app_command( "dump_chain" => DumpChainCommand::handle(inter, data, ctx).await, "load_chain" => LoadChainCommand::handle(inter, data, ctx).await, "markov" => MarkovCommand::handle(inter, data, ctx).await, + "forget" => ForgetCommand::handle(inter, data, ctx).await, other => { warn!("Unknown command send: {other}"); Ok(())