diff --git a/src/bot/commands.rs b/src/bot/commands.rs index cd65995..2b26c1d 100644 --- a/src/bot/commands.rs +++ b/src/bot/commands.rs @@ -86,7 +86,15 @@ impl Bot { prefix: &self.config.command_prefix, network: &self.config.name, }; - let actions = self.modules[idx].on_command(&cmd); + let actions = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + self.modules[idx].on_command(&cmd) + })) { + Ok(actions) => actions, + Err(_) => { + crate::log(&format!("module panicked on '{command}'")); + return Ok(()); + } + }; for action in actions { match action { Action::Reply(text) => self.conn.privmsg(reply_to, &text)?, diff --git a/src/bot/modules/weather.rs b/src/bot/modules/weather.rs index 4acbe8d..d8b5c06 100644 --- a/src/bot/modules/weather.rs +++ b/src/bot/modules/weather.rs @@ -1,16 +1,18 @@ -use std::collections::BTreeMap; +use std::collections::{BTreeMap, HashMap}; use std::fs; use std::io::{self, Read, Write}; use std::iter::Peekable; use std::net::{TcpStream, ToSocketAddrs}; use std::path::{Path, PathBuf}; use std::str::Chars; -use std::time::Duration; +use std::time::{Duration, Instant}; use native_tls::TlsConnector; use crate::bot::module::{Action, Command, CommandSpec, Module}; +const CACHE_TTL: Duration = Duration::from_secs(600); + const COMMANDS: &[CommandSpec] = &[ CommandSpec { name: "add", usage: "add ", about: "save your weather location" }, CommandSpec { name: "w", usage: "w [location]", about: "weather for your saved location" }, @@ -19,6 +21,7 @@ const COMMANDS: &[CommandSpec] = &[ pub struct Weather { store: Store, path: PathBuf, + cache: HashMap, } impl Weather { @@ -32,7 +35,7 @@ impl Weather { } pub fn with_path(path: PathBuf) -> Weather { - Weather { store: Store::load(&path), path } + Weather { store: Store::load(&path), path, cache: HashMap::new() } } pub fn location_of(&self, nick: &str) -> Option<&str> { @@ -74,8 +77,17 @@ impl Module for Weather { cmd.sender, cmd.prefix ))]; }; + let key = loc.to_lowercase(); + if let Some((at, text)) = self.cache.get(&key) { + if at.elapsed() < CACHE_TTL { + return vec![Action::reply(text.clone())]; + } + } match fetch_weather(&loc) { - Ok(text) => vec![Action::reply(text)], + Ok(text) => { + self.cache.insert(key, (Instant::now(), text.clone())); + vec![Action::reply(text)] + } Err(e) => vec![Action::reply(format!("weather unavailable: {e}"))], } } diff --git a/src/config.rs b/src/config.rs index 1c7011c..af5289e 100644 --- a/src/config.rs +++ b/src/config.rs @@ -14,6 +14,7 @@ pub struct Config { pub realname: String, pub channels: Vec, pub command_prefix: String, + pub verbose: bool, pub password: Option, pub sasl_user: Option, pub sasl_pass: Option, @@ -31,6 +32,7 @@ impl Default for Config { realname: "rubot".to_string(), channels: vec!["#devs".to_string()], command_prefix: "!".to_string(), + verbose: true, password: None, sasl_user: None, sasl_pass: None, @@ -146,6 +148,7 @@ fn apply_kv(c: &mut Config, key: &str, value: &str, where_: &str) { "realname" => c.realname = value.to_string(), "channels" => c.channels = split_list(value), "prefix" | "command_prefix" => c.command_prefix = value.to_string(), + "verbose" => c.verbose = parse_bool(value), "password" => c.password = (!value.is_empty()).then(|| value.to_string()), "sasl_user" => c.sasl_user = (!value.is_empty()).then(|| value.to_string()), "sasl_pass" => c.sasl_pass = (!value.is_empty()).then(|| value.to_string()), @@ -194,6 +197,9 @@ fn apply_env(c: &mut Config) { if let Ok(v) = std::env::var("IRC_PREFIX") { c.command_prefix = v; } + if let Ok(v) = std::env::var("IRC_VERBOSE") { + c.verbose = parse_bool(&v); + } if let Ok(v) = std::env::var("IRC_PASSWORD") { c.password = Some(v); } diff --git a/src/irc/connection.rs b/src/irc/connection.rs index 8c46ba6..ada8eda 100644 --- a/src/irc/connection.rs +++ b/src/irc/connection.rs @@ -1,20 +1,27 @@ use std::io::{self, BufRead, BufReader, Read, Write}; use std::net::TcpStream; +use std::time::{Duration, Instant}; use native_tls::TlsConnector; use super::Message; +const MAX_LINE: usize = 500; +const SEND_THROTTLE: Duration = Duration::from_millis(300); + trait Stream: Read + Write {} impl Stream for T {} pub struct Connection { inner: BufReader>, log_prefix: String, + verbose: bool, + throttle: Duration, + last_send: Option, } impl Connection { - pub fn connect(host: &str, port: u16, use_tls: bool) -> io::Result { + pub fn connect(host: &str, port: u16, use_tls: bool, verbose: bool) -> io::Result { let tcp = TcpStream::connect((host, port))?; let stream: Box = if use_tls { let connector = @@ -26,7 +33,24 @@ impl Connection { } else { Box::new(tcp) }; - Ok(Connection { inner: BufReader::new(stream), log_prefix: String::new() }) + Ok(Connection { + inner: BufReader::new(stream), + log_prefix: String::new(), + verbose, + throttle: SEND_THROTTLE, + last_send: None, + }) + } + + pub fn from_stream(stream: S) -> Connection { + let boxed: Box = Box::new(stream); + Connection { + inner: BufReader::new(boxed), + log_prefix: String::new(), + verbose: false, + throttle: Duration::ZERO, + last_send: None, + } } pub fn set_label(&mut self, label: &str) { @@ -47,7 +71,9 @@ impl Connection { if trimmed.is_empty() { continue; } - eprintln!("{}<< {trimmed}", self.log_prefix); + if self.verbose { + eprintln!("{}<< {trimmed}", self.log_prefix); + } if let Some(msg) = Message::parse(trimmed) { return Ok(Some(msg)); } @@ -55,8 +81,20 @@ impl Connection { } pub fn send_raw(&mut self, line: &str) -> io::Result<()> { + if !self.throttle.is_zero() { + if let Some(last) = self.last_send { + let elapsed = last.elapsed(); + if elapsed < self.throttle { + std::thread::sleep(self.throttle - elapsed); + } + } + self.last_send = Some(Instant::now()); + } let clean = line.replace(['\r', '\n'], " "); - eprintln!("{}>> {clean}", self.log_prefix); + let clean = truncate_to(&clean, MAX_LINE); + if self.verbose { + eprintln!("{}>> {clean}", self.log_prefix); + } let writer = self.inner.get_mut(); write!(writer, "{clean}\r\n")?; writer.flush() @@ -110,3 +148,14 @@ impl Connection { self.send_raw(&format!("AUTHENTICATE {payload}")) } } + +fn truncate_to(s: &str, max: usize) -> &str { + if s.len() <= max { + return s; + } + let mut end = max; + while end > 0 && !s.is_char_boundary(end) { + end -= 1; + } + &s[..end] +} diff --git a/src/lib.rs b/src/lib.rs index 70bc600..4926187 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,5 @@ +#![forbid(unsafe_code)] + pub mod bot; pub mod config; pub mod irc; diff --git a/src/main.rs b/src/main.rs index 2782d81..f85853b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -61,7 +61,7 @@ fn supervise(config: Config, labelled: bool) { } fn run_once(config: &Config, labelled: bool) -> std::io::Result<()> { - let mut conn = Connection::connect(&config.server, config.port, config.tls)?; + let mut conn = Connection::connect(&config.server, config.port, config.tls, config.verbose)?; if labelled { conn.set_label(&config.name); } diff --git a/tests/bot.rs b/tests/bot.rs new file mode 100644 index 0000000..fd6d040 --- /dev/null +++ b/tests/bot.rs @@ -0,0 +1,78 @@ +use std::io::{self, Cursor, Read, Write}; +use std::sync::{Arc, Mutex}; + +use rustbot::bot::Bot; +use rustbot::config::Config; +use rustbot::irc::Connection; + +struct Mock { + input: Cursor>, + output: Arc>>, +} + +impl Read for Mock { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.input.read(buf) + } +} + +impl Write for Mock { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.output.lock().unwrap().extend_from_slice(buf); + Ok(buf.len()) + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +fn run(server: &str) -> String { + let output = Arc::new(Mutex::new(Vec::new())); + let mock = Mock { + input: Cursor::new(server.as_bytes().to_vec()), + output: output.clone(), + }; + let conn = Connection::from_stream(mock); + let mut bot = Bot::new(Config::default(), conn); + let _ = bot.run(); + let bytes = output.lock().unwrap().clone(); + String::from_utf8_lossy(&bytes).into_owned() +} + +#[test] +fn sends_registration() { + let out = run(""); + assert!(out.contains("CAP LS"), "{out}"); + assert!(out.contains("NICK rubot"), "{out}"); + assert!(out.contains("USER rubot"), "{out}"); +} + +#[test] +fn answers_ping() { + let out = run("PING :tok123\r\n"); + assert!(out.contains("PONG :tok123"), "{out}"); +} + +#[test] +fn joins_configured_channels_on_welcome() { + let out = run(":srv 001 rubot :welcome\r\n"); + assert!(out.contains("JOIN #devs"), "{out}"); +} + +#[test] +fn appends_underscore_on_nick_in_use() { + let out = run(":srv 433 * rubot :nickname is already in use\r\n"); + assert!(out.contains("NICK rubot_"), "{out}"); +} + +#[test] +fn dispatches_command_to_module() { + let out = run(":alice!a@h PRIVMSG #chan :!ping\r\n"); + assert!(out.contains("PRIVMSG #chan :pong"), "{out}"); +} + +#[test] +fn ignores_own_echoed_message() { + let out = run(":rubot!r@h PRIVMSG #chan :!ping\r\n"); + assert!(!out.contains(":pong"), "{out}"); +}