diff --git a/.gitignore b/.gitignore index ede43cf..52d9946 100644 --- a/.gitignore +++ b/.gitignore @@ -5,4 +5,7 @@ Cargo.lock # Logs -*.log \ No newline at end of file +*.log + +# Ignore db +database.db* \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..8a96a12 --- /dev/null +++ b/README.md @@ -0,0 +1,10 @@ +# Chat Client and Server +a very simple IRC clone + + +## What does it do? +The client connects to the server and sends messages to the server. The server then broadcasts the message to all connected clients. + +The client can choose a username and a "Group" or channel where the message will be send. This allows for multiple conversations to be happening at the same time on the same server. New clients can join the group and will receive all messages sent to that group. + +Old messages are stored on the server in a sqlite database. \ No newline at end of file diff --git a/client/src/main.rs b/client/src/main.rs index bd30413..709da97 100644 --- a/client/src/main.rs +++ b/client/src/main.rs @@ -39,7 +39,7 @@ async fn main() { // Run until the application returns false while Application::new(address, name, group).run().await { - error!("Application Disconnected. Press any key to reconnect"); + error!("Server Disconnected, press enter to try to reconnect"); } TUI::exit().expect("Failed to reset terminal"); diff --git a/client/src/tui.rs b/client/src/tui.rs index 4d0bc67..b69af52 100644 --- a/client/src/tui.rs +++ b/client/src/tui.rs @@ -92,7 +92,10 @@ impl TUI { for message in messages { let line = Line::from(vec![ if message.username == model.username { - Span::styled(&message.username, ratatui::style::Style::default().bold()) + Span::styled( + &message.username, + ratatui::style::Style::default().bold().on_dark_gray(), + ) } else { Span::styled(&message.username, ratatui::style::Style::default()) }, diff --git a/server/Cargo.toml b/server/Cargo.toml index 7fe2ea6..c691743 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -17,3 +17,6 @@ anyhow = "1.0.76" # Logging log = "0.4.20" env_logger = "0.10.1" + +# Database +sqlx = { version = "0.7.4", features = ["sqlite", "runtime-tokio", "macros"] } diff --git a/server/src/app.rs b/server/src/app.rs index 4cbad77..c80a05c 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -4,15 +4,18 @@ use std::{ }; use futures_util::FutureExt; -use log::{debug, error}; +use log::{debug, error, info}; -use crate::{connection::Connection, websocket}; +use crate::{connection::Connection, database, websocket}; type Sender = tokio::sync::mpsc::UnboundedSender; +static DATABASE_URL: &str = "sqlite://database.db"; + pub struct Application { pub adress: String, pub connections: Arc>>>, + pub db: Option, } impl Application { @@ -20,10 +23,26 @@ impl Application { Self { adress: adress.to_string(), connections: Arc::new(Mutex::new(HashMap::new())), + db: None, } } pub async fn run(&mut self) { + // Connect to DB + if self.db.is_none() { + match database::establish_connection(DATABASE_URL).await { + Ok(db) => { + self.db = Some(db); + info!("Connected to database at {}", DATABASE_URL); + } + Err(e) => { + error!("Failed to connect to database: {}", e); + } + } + + database::create_table(self.db.as_ref().unwrap()).await; + } + let (connection_sender, mut connection_receiver) = tokio::sync::mpsc::unbounded_channel(); //let (message_sender, message_receiver) = tokio::sync::mpsc::unbounded_channel(); @@ -49,17 +68,32 @@ impl Application { let mut read_channel = connection.receiver; let connections = self.connections.clone(); + let db_connection = self.db.as_ref().cloned(); + tokio::spawn(async move { // await for the group to be set while connection.group.lock().unwrap().is_none() { tokio::task::yield_now().await; } + let group = connection.group.lock().unwrap().as_ref().unwrap().clone(); + + // Send all messages from the database + if let Some(ref db) = db_connection { + let messages = crate::database::get_messages(db, &group).await; + info!("Sending {} messages from group '{}'", messages.len(), group); + + for message in messages { + if let Err(e) = Connection::send(&connection.sender, message.serialize()) { + error!("Error sending message: {}", e); + } + } + } - // Add the connection to the group + // Add the connection to the list of connections connections .lock() .unwrap() - .entry(connection.group.lock().unwrap().as_ref().unwrap().clone()) + .entry(group.clone()) .or_default() .push(connection.sender); @@ -68,15 +102,35 @@ impl Application { let msg = read_channel.recv().await; if let Some(msg) = msg { debug!("Recieved message: {}", msg); - let mut connections = connections.lock().unwrap(); + let parsed = match crate::ChatMessage::deserialize(&msg) { + Some(msg) => { + debug!("Message from {}: {}", msg.username, msg.message); + msg + } + None => { + error!("Failed to deserialize message: {}", msg); + continue; + } + }; + + // Save message to database + if let Some(ref db) = db_connection { + crate::database::insert_message(db, &group, &parsed).await; + } + + // Send message to everyone in the group + let mut connections = connections.lock().unwrap(); if let Some(group) = connection.group.lock().unwrap().as_ref() { if let Some(connections) = connections.get_mut(group) { - for connection in connections.iter_mut() { - if let Err(e) = Connection::send(connection, msg.clone()) { + connections.retain(|c| { + if let Err(e) = Connection::send(c, msg.clone()) { error!("Error sending message: {}", e); + false + } else { + true } - } + }); } } } else { diff --git a/server/src/database.rs b/server/src/database.rs new file mode 100644 index 0000000..b359159 --- /dev/null +++ b/server/src/database.rs @@ -0,0 +1,81 @@ +use log::info; +use sqlx::{migrate::MigrateDatabase, sqlite::SqlitePoolOptions, Pool, Sqlite}; + +use crate::ChatMessage; + +static MESSAGE_RETRIVAL_AMOUNT: u32 = 100; + +pub async fn establish_connection(database_url: &str) -> anyhow::Result> { + // Create database if needed + if !Sqlite::database_exists(database_url).await.unwrap_or(false) { + info!("Database does not exist, creating it"); + Sqlite::create_database(database_url).await?; + } + + SqlitePoolOptions::new() + .max_connections(10) + .connect(database_url) + .await + .map_err(|e| e.into()) +} + +pub async fn create_table(pool: &Pool) { + sqlx::query( + r#" + CREATE TABLE IF NOT EXISTS messages ( + id INTEGER PRIMARY KEY, + group_name TEXT NOT NULL, + username TEXT NOT NULL, + message TEXT NOT NULL + ) + "#, + ) + .execute(pool) + .await + .expect("Failed to create table"); + + // Count number of messages + let count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM messages") + .fetch_one(pool) + .await + .expect("Failed to count messages"); + info!("Database contains {} messages", count); +} + +pub async fn insert_message(pool: &Pool, group_name: &str, message: &ChatMessage) { + sqlx::query( + r#" + INSERT INTO messages (group_name, username, message) + VALUES (?, ?, ?) + "#, + ) + .bind(group_name) + .bind(&message.username) + .bind(&message.message) + .execute(pool) + .await + .expect("Failed to insert message"); +} + +pub async fn get_messages(pool: &Pool, group_name: &str) -> Vec { + sqlx::query_as( + r#" + SELECT username, message + FROM messages + WHERE group_name = ? + ORDER BY id ASC + LIMIT ? + "#, + ) + .bind(group_name) + .bind(MESSAGE_RETRIVAL_AMOUNT) + .fetch_all(pool) + .await + .map(|messages: Vec<(String, String)>| { + messages + .into_iter() + .map(|(username, message)| ChatMessage { username, message }) + .collect() + }) + .expect("Failed to fetch messages") +} diff --git a/server/src/lib.rs b/server/src/lib.rs index 539f85d..b008ccb 100644 --- a/server/src/lib.rs +++ b/server/src/lib.rs @@ -1,3 +1,23 @@ -pub mod websocket; +pub mod app; pub mod connection; -pub mod app; \ No newline at end of file +pub mod database; +pub mod websocket; + +pub struct ChatMessage { + pub username: String, + pub message: String, +} + +impl ChatMessage { + pub fn serialize(&self) -> String { + format!("{}: {}", self.username, self.message) + } + + // Security issue: a username can contain ": ", which would break the deserialization + pub fn deserialize(s: &str) -> Option { + let mut parts = s.splitn(2, ": "); + let username = parts.next()?.to_string(); + let message = parts.next()?.to_string(); + Some(Self { username, message }) + } +}