blatherskite

a toy discord-like chat app backend written for a swe class
Log | Files | Refs | README

commit b29bbcb99286d26ec98974c6d2b66767aed5ef0e
parent 5a988fd14016aaa2db56cb6d00db7ca132d3c38b
Author: quantumish <freifeld.david@gmail.com>
Date:   Sat, 15 Oct 2022 22:59:30 -0700

Major simplification of `scuttlebutt` code

Diffstat:
Mscuttlebutt/Cargo.toml | 2++
Ascuttlebutt/build.rs | 6++++++
Ascuttlebutt/src/db.rs | 344+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Mscuttlebutt/src/main.rs | 655++++++++++++++++++-------------------------------------------------------------
Mscuttlebutt/src/tests.rs | 202+++++++++++++++++++++++++++++++++++++++++--------------------------------------
5 files changed, 602 insertions(+), 607 deletions(-)

diff --git a/scuttlebutt/Cargo.toml b/scuttlebutt/Cargo.toml @@ -11,6 +11,7 @@ base64 = "0.13.0" cassandra-cpp = "1.1.0" chrono = { version = "0.4.22", features = ["serde"] } ctor = "0.1.23" +error-stack = "0.2.3" hex = "0.4.3" hmac = "0.12.1" jwt = "0.16.0" @@ -24,5 +25,6 @@ rustflake = "0.1.1" serde = "1.0.144" serde_json = "1.0.85" sha2 = "0.10.6" +thiserror = "1.0.37" tokio = { version = "1", features = ["full"] } tracing-subscriber = "0.3.15" diff --git a/scuttlebutt/build.rs b/scuttlebutt/build.rs @@ -0,0 +1,6 @@ +fn main() { + if cfg!(target_os = "macos") { + println!("cargo:rustc-link-search=native=/opt/local/lib"); + println!("cargo:rustc-link-search=native=/opt/homebrew/lib"); + } +} diff --git a/scuttlebutt/src/db.rs b/scuttlebutt/src/db.rs @@ -0,0 +1,344 @@ +use cassandra_cpp::*; +use crate::responses::*; + +#[derive(Debug)] +pub enum IdType { + User, + Group, + Channel, + Message +} + +pub trait Database: Sync + Send { + fn valid_id(&self, kind: IdType, id: i64) -> Result<bool>; + + fn create_user(&self, id: i64, name: String, email: String, hash: String) -> Result<()>; + fn update_user(&self, id: i64, name: String, email: String) -> Result<()>; + fn get_user(&self, id: i64) -> Result<User>; + fn get_user_hash(&self, id: i64) -> Result<String>; + fn delete_user(&self, id: i64) -> Result<()>; + + fn create_group(&self, gid: i64, uid: i64, name: String) -> Result<()>; + fn update_group(&self, id: i64, name: String) -> Result<()>; + fn get_group(&self, id: i64) -> Result<Group>; + fn delete_group(&self, id: i64) -> Result<()>; + fn get_group_members(&self, gid: i64) -> Result<Vec<i64>>; + fn remove_group_member(&self, gid: i64, uid: i64) -> Result<()>; + fn add_group_member(&self, gid: i64, uid: i64) -> Result<()>; + fn get_group_channels(&self, gid: i64) -> Result<Vec<i64>>; + fn remove_group_channel(&self, gid: i64, uid: i64) -> Result<()>; + fn add_group_channel(&self, gid: i64, uid: i64) -> Result<()>; + + fn create_channel(&self, cid: i64, gid: i64, uid: i64, name: String) -> Result<()>; + fn get_channel(&self, id: i64) -> Result<Channel>; + fn update_channel(&self, id: i64, name: String) -> Result<()>; + fn delete_channel(&self, id: i64) -> Result<()>; + fn get_channel_members(&self, gid: i64) -> Result<Vec<i64>>; + fn remove_channel_member(&self, cid: i64, id: i64) -> Result<()>; + fn add_channel_member(&self,cid: i64, id: i64) -> Result<()>; + + fn get_user_groups(&self, id: i64) -> Result<Vec<i64>>; + fn delete_user_groups(&self, id: i64) -> Result<()>; + fn add_user_group(&self, uid: i64, gid: i64) -> Result<()>; + fn remove_user_group(&self, uid: i64, gid: i64) -> Result<()>; + + fn get_message(&self, id: i64) -> Result<Message>; + fn get_messages(&self, cid: i64, num: u64) -> Result<Vec<Message>>; +} + +pub struct Cassandra { + kspc: String, + sess: Session +} + +impl Cassandra { + pub fn new(keyspc: &str) -> Self { + let contact_points = "127.0.0.1"; + let mut cluster = Cluster::default(); + cluster.set_contact_points(contact_points).unwrap(); + cluster.set_load_balance_round_robin(); + let session = cluster.connect().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE KEYSPACE IF NOT EXISTS {keyspc} \ + WITH replication = {{'class':'SimpleStrategy', 'replication_factor': 1}}" + ))).wait().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE TABLE IF NOT EXISTS {keyspc}.users \ + (id bigint PRIMARY KEY, name text, email text, hash text);" + ))).wait().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE TABLE IF NOT EXISTS {keyspc}.groups \ + (id bigint PRIMARY KEY, name text, members set<bigint>, channels set<bigint>);" + ))).wait().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE TABLE IF NOT EXISTS {keyspc}.channels \ + (id bigint PRIMARY KEY, name text, group bigint, members set<bigint>);" + ))).wait().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE TABLE IF NOT EXISTS {keyspc}.user_groups \ + (id bigint PRIMARY KEY, groups set<bigint>);" + ))).wait().unwrap(); + + session.execute(&stmt!(&format!( + "CREATE TABLE IF NOT EXISTS {keyspc}.messages \ + (channel bigint, id bigint, author bigint, \ + time timestamp, content text, PRIMARY KEY (channel, id)) \ + WITH CLUSTERING ORDER BY (id DESC);" + ))).wait().unwrap(); + + Self { + kspc: keyspc.to_string(), + sess: session + } + } + + fn delete_row(&self, table: &str, id: i64) -> Result<()> { + self.sess.execute(&stmt!(&format!( + "DELETE FROM {}.{table} WHERE id={id};", self.kspc + ))).wait().unwrap(); + Ok(()) + } + + fn get_set(&self, table: &str, set: &str, id: i64) -> Result<Vec<i64>> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT {set} FROM {}.{table} WHERE id = {id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + let items: SetIterator = row.get(0)?; + Ok(items.map(|i| i.get_i64().unwrap()).collect()) + } + + fn pop_set(&self, table: &str, set: &str, id: i64, elem: i64) -> Result<()> { + self.sess.execute(&stmt!(&format!( + "UPDATE {}.{table} SET {set} = {set} - {{{elem}}} WHERE ID={id};", self.kspc + ))).wait()?; + Ok(()) + } + + fn push_set(&self, table: &str, set: &str, id: i64, elem: i64) -> Result<()> { + self.sess.execute(&stmt!(&format!( + "UPDATE {}.{table} SET {set} = {set} + {{{elem}}} WHERE ID={id};", self.kspc + ))).wait()?; + Ok(()) + } +} + +impl Database for Cassandra { + fn valid_id(&self, kind: IdType, id: i64) -> Result<bool> { + let table = match kind { + IdType::User => "users", + IdType::Group => "groups", + IdType::Channel => "channels", + IdType::Message => "messages", + }; + let res = self.sess.execute(&stmt!(&format!( + "SELECT * FROM {}.{table} WHERE ID={id};", self.kspc + ))).wait()?; + if let Some(_row) = res.first_row() { + return Ok(true) + } else { return Ok(false) }; + } + + fn create_user(&self, id: i64, name: String, email: String, hash: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "INSERT INTO {}.users (id, name, email, hash) VALUES ({id}, ?, ?, ?);", self.kspc + )); + stmt.bind(0, name.as_str())?; + stmt.bind(1, email.as_str())?; + stmt.bind(2, hash.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn get_user(&self, id: i64) -> Result<User> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT name, email FROM {}.users WHERE ID={id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + Ok(User { + id, + username: row.get(0)?, + email: row.get(1)? + }) + } + + fn get_user_hash(&self, id: i64) -> Result<String> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT hash FROM {}.users WHERE ID={id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + Ok(row.get(0)?) + } + + fn update_user(&self, id: i64, name: String, email: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "UPDATE {}.users SET name=?, email=? WHERE ID={id};", self.kspc + )); + stmt.bind(0, name.as_str())?; + stmt.bind(1, email.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn delete_user(&self, id: i64) -> Result<()> { + self.delete_row("users", id) + } + + fn get_group(&self, id: i64) -> Result<Group> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT name, members, channels FROM {}.groups WHERE ID={id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + let members: SetIterator = row.get(1)?; + let channels: SetIterator = row.get(2)?; + Ok(Group { + id, + name: row.get(0)?, + members: members.map(|i| i.get_i64().unwrap()).collect(), + channels: channels.map(|i| i.get_i64().unwrap()).collect(), + }) + } + + fn create_group(&self, gid: i64, uid: i64, name: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "INSERT INTO {}.groups (id, name, channels, members) VALUES ({gid}, ?, {{}}, {{{uid}}});", self.kspc + )); + stmt.bind(0, name.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn delete_group(&self, id: i64) -> Result<()> { + self.delete_row("groups", id) + } + + fn update_group(&self, id: i64, name: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "UPDATE {}.groups SET name = ? WHERE id = {id};", self.kspc + )); + stmt.bind(0, name.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn get_group_members(&self, gid: i64) -> Result<Vec<i64>> { + self.get_set("groups", "members", gid) + } + + fn get_group_channels(&self, gid: i64) -> Result<Vec<i64>> { + self.get_set("groups", "channels", gid) + } + + fn add_group_member(&self, gid: i64, uid: i64) -> Result<()> { + self.push_set("groups", "members", gid, uid) + } + + fn remove_group_member(&self, gid: i64, uid: i64) -> Result<()> { + self.pop_set("groups", "members", gid, uid) + } + + fn add_group_channel(&self, gid: i64, uid: i64) -> Result<()> { + self.push_set("groups", "channels", gid, uid) + } + + fn remove_group_channel(&self, gid: i64, uid: i64) -> Result<()> { + self.pop_set("groups", "channels", gid, uid) + } + + fn get_channel(&self, id: i64) -> Result<Channel> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT group, name, members FROM {}.channels WHERE ID={id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + let members: SetIterator = row.get(2)?; + Ok(Channel { + id, + group: row.get(0)?, + name: row.get(1)?, + members: members.map(|i| i.get_i64().unwrap()).collect(), + }) + } + + fn create_channel(&self, cid: i64, gid: i64, uid: i64, name: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "INSERT INTO {}.channels (id, group, name, members) VALUES ({cid}, {gid}, ?, {{{uid}}});", self.kspc + )); + stmt.bind(0, name.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn delete_channel(&self, id: i64) -> Result<()> { + self.delete_row("channels", id) + } + + fn update_channel(&self, id: i64, name: String) -> Result<()> { + let mut stmt = stmt!(&format!( + "UPDATE {}.channels SET name = ? WHERE id = {id};", self.kspc + )); + stmt.bind(0, name.as_str())?; + self.sess.execute(&stmt).wait()?; + Ok(()) + } + + fn get_channel_members(&self, cid: i64) -> Result<Vec<i64>> { + self.get_set("channel", "members", cid) + } + + fn add_channel_member(&self, gid: i64, uid: i64) -> Result<()> { + self.push_set("channels", "members", gid, uid) + } + + fn remove_channel_member(&self, gid: i64, uid: i64) -> Result<()> { + self.pop_set("channels", "members", gid, uid) + } + + fn get_user_groups(&self, id: i64) -> Result<Vec<i64>> { + self.get_set("user_groups", "groups", id) + } + + fn add_user_group(&self, uid: i64, gid: i64) -> Result<()> { + self.push_set("user_groups", "groups", uid, gid) + } + + fn remove_user_group(&self, uid: i64, gid: i64) -> Result<()> { + self.pop_set("user_groups", "groups", uid, gid) + } + + fn delete_user_groups(&self, id: i64) -> Result<()> { + self.delete_row("user_groups", id) + } + + // TODO the unwraps here are not great + fn get_messages(&self, cid: i64, num: u64) -> Result<Vec<Message>> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT * FROM {}.messages WHERE channel={cid} LIMIT {num};", self.kspc + ))).wait()?; + Ok(res.iter().map(|row| { + Message { + id: row.get(0).unwrap(), + author: row.get(2).unwrap(), + channel: row.get(1).unwrap(), + content: row.get(4).unwrap() + } + }).collect::<Vec<Message>>()) + } + + fn get_message(&self, id: i64) -> Result<Message> { + let res = self.sess.execute(&stmt!(&format!( + "SELECT channel, author, content, time FROM {}.users WHERE ID={id};", self.kspc + ))).wait()?; + let row = res.first_row().unwrap(); + Ok(Message { + id, + channel: row.get(0)?, + author: row.get(1)?, + content: row.get(2)?, + }) + } +} diff --git a/scuttlebutt/src/main.rs b/scuttlebutt/src/main.rs @@ -1,9 +1,8 @@ -use cassandra_cpp::*; -use chrono::{DateTime, Duration, Local}; +use chrono::{DateTime, Duration, Local, Utc}; use hmac::Hmac; use jwt::{SignWithKey, VerifyWithKey}; use poem::{ - http::StatusCode, listener::TcpListener, web::Data, Endpoint, EndpointExt, Request, Result, + listener::TcpListener, web::Data, EndpointExt, Request, Result, Route, Server, }; use poem_openapi::{ @@ -16,14 +15,14 @@ use rand::{distributions::Alphanumeric, Rng}; use rustflake::Snowflake; use serde::{Deserialize, Serialize}; use sha2::Sha256; -use std::sync::Mutex; pub mod responses; pub use responses::*; -type ServerKey = Hmac<Sha256>; +pub mod db; +pub use db::*; -const UNUSUAL_ROW_ERROR: &'static str = "Found duplicate ID (or negative rows??)! Giving up!"; +type ServerKey = Hmac<Sha256>; #[derive(Serialize, Deserialize)] struct Claims { @@ -55,113 +54,54 @@ async fn api_checker(req: &Request, api_key: ApiKey) -> Option<Claims> { } struct Api { - sess: Session, - kspc: String, + db: Box<dyn Database>, } pub fn gen_id() -> i64 { - static STATE: Mutex<Option<Snowflake>> = Mutex::new(None); - - STATE - .lock() - .unwrap() - .get_or_insert_with(|| Snowflake::default()) - .generate() + // Very cursed thread-unique number generation + // NOTE: substitute for std::thread::current().id() when it's stabilized + thread_local! { static V: u8 = 0; } + let id: u64 = V.with(|v| v as *const u8 as u64); + + let now = Utc::now().timestamp_nanos(); + // TODO generalize this hardcoded 1 + Snowflake::new(now, 1, id as i64).generate() } #[OpenApi] #[allow(unused_variables)] impl Api { - fn new(keyspc: &str) -> Api { - let contact_points = "127.0.0.1"; - let mut cluster = Cluster::default(); - cluster.set_contact_points(contact_points).unwrap(); - cluster.set_load_balance_round_robin(); - let session = cluster.connect().unwrap(); - - session.execute(&stmt!(&format!("CREATE KEYSPACE IF NOT EXISTS {keyspc} WITH replication = {{'class':'SimpleStrategy', 'replication_factor': 1}}"))).wait().unwrap(); - session.execute(&stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.users (id bigint PRIMARY KEY, name text, email text, hash text);"))).wait().unwrap(); - session.execute(&stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.groups (id bigint PRIMARY KEY, name text, members set<bigint>, channels set<bigint>);"))).wait().unwrap(); - session.execute(&stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.channels (id bigint PRIMARY KEY, name text, group bigint, members set<bigint>);"))).wait().unwrap(); - session.execute(&stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.user_groups (id bigint PRIMARY KEY, groups set<bigint>);"))).wait().unwrap(); - session.execute(&stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.messages (channel bigint, id bigint, author bigint, time timestamp, content text, PRIMARY KEY (channel, id)) WITH CLUSTERING ORDER BY (id DESC);"))).wait().unwrap(); - - Api { - sess: session, - kspc: String::from(keyspc), - } - } - - fn validate_id(&self, table: &str, gid: i64) -> anyhow::Result<()> { - let res = self.sess.execute(&stmt!(&format!( - "SELECT id FROM {}.groups WHERE id = {};", self.kspc, gid - ))).wait().unwrap(); - match res.row_count() { - 1 => Ok(()), - _ => Err(anyhow::anyhow!("not found")), - } - - } - - fn __remove_channel_member(&self, cid: i64, uid: i64) { - self.sess.execute(&stmt!(&format!( - "UPDATE {}.channels SET members = members - {{{}}} WHERE id={};", self.kspc, uid, cid - ))).wait().unwrap(); + fn new(db: Box<dyn Database>) -> Api { + Api { db } } fn __remove_group_member(&self, gid: i64, uid: i64) { - self.sess.execute(&stmt!(&format!( - "UPDATE {}.groups SET members = members - {{{}}} WHERE id={};", self.kspc, uid, gid - ))).wait().unwrap(); - let res = self.sess.execute(&stmt!(&format!( - "SELECT channels FROM {}.groups WHERE id={};", self.kspc, gid - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let channels: SetIterator = row.get(0).unwrap(); + self.db.remove_group_member(gid, uid).unwrap(); + let channels = self.db.get_group_channels(gid).unwrap(); for channel in channels { - self.__remove_channel_member(channel.get_i64().unwrap(), uid); + self.db.remove_channel_member(channel, uid).unwrap(); } - self.sess.execute(&stmt!(&format!( - "UPDATE {}.user_groups SET groups = groups - {{{}}} WHERE id={};", self.kspc, gid, uid - ))).wait().unwrap(); + self.db.remove_user_group(uid, gid).unwrap(); } #[oai(path = "/login", method = "post")] - async fn login( - &self, - key: Data<&ServerKey>, - id: Query<i64>, - hash: PlainText<String>, - ) -> LoginResponse { + async fn login(&self, key: Data<&ServerKey>, id: Query<i64>, hash: PlainText<String>) -> LoginResponse { use LoginResponse::*; if hash.0.len() != 64 { return BadRequest; + } else if !self.db.valid_id(IdType::User, id.0).unwrap() { + return NotFound; } - let hash_stmt = &stmt!(&format!( - "SELECT id, name, email, hash FROM {}.users WHERE id={};", - self.kspc, id.0 - )); - let res = self.sess.execute(hash_stmt).wait().unwrap(); - if res.row_count() == 1 { - let row = res.first_row().unwrap(); - let db_hash: String = row.get(3).unwrap(); - if hex::decode(db_hash).unwrap() != hex::decode(hash.0).unwrap() { - Unauthorized - } else { - let row = res.first_row().unwrap(); - let token = Claims { - id: id.0, - exp: Local::now() + Duration::days(1), - } - .sign_with_key(key.0); - Success(PlainText(token.unwrap())) - } - } else if res.row_count() == 0 { - NotFound + let db_hash = self.db.get_user_hash(id.0).unwrap(); + if hex::decode(db_hash).unwrap() != hex::decode(hash.0).unwrap() { + Unauthorized } else { - InternalError(PlainText( - "Found multiple (or negative?) number of rows.".to_string(), - )) + let token = Claims { + id: id.0, + exp: Local::now() + Duration::days(1), + } + .sign_with_key(key.0); + Success(PlainText(token.unwrap())) } } @@ -173,104 +113,46 @@ impl Api { /// Call `/user?id=1234` to get the user with id 1234 async fn get_user(&self, id: Query<i64>) -> UserResponse { use UserResponse::*; - let insert_stmt = &stmt!(&format!( - "SELECT id, name, email FROM {}.users WHERE id={};", - self.kspc, id.0 - )); - match self.sess.execute(insert_stmt).wait() { - Ok(res) => { - if res.row_count() == 1 { - let row = res.first_row().unwrap(); - Success(Json(User { - id: id.0, - username: row.get(1).unwrap(), - email: row.get(2).unwrap(), - })) - } else if res.row_count() == 0 { - NotFound - } else { - InternalError(PlainText(UNUSUAL_ROW_ERROR.to_string())) - } - } - Err(e) => InternalError(PlainText(e.to_string())), + if !self.db.valid_id(IdType::User, id.0).unwrap() { return NotFound; } + match self.db.get_user(id.0) { + Ok(user) => Success(Json(user)), + Err(e) => InternalError(PlainText(e.to_string())) } } #[oai(path = "/user", method = "post")] /// Creates a new user - async fn make_user( - &self, - name: Query<String>, - email: Query<String>, - hash: Query<String>, - ) -> CreateUserResponse { + async fn make_user(&self, name: Query<String>, email: Query<String>, hash: Query<String>) -> CreateUserResponse { use CreateUserResponse::*; if hash.0.len() != 64 { return BadRequest(PlainText("Invalid hash provided.".to_string())); } let id = gen_id(); - let insert_stmt = &stmt!(&format!( - "INSERT INTO {}.users (id, name, email, hash) VALUES ({},'{}','{}','{}');", - self.kspc, id, name.0, email.0, hash.0 - )); - if let Err(e) = self.sess.execute(insert_stmt).wait() { - InternalError(PlainText(e.to_string())) - } else { - Success(Json(User { - id, - username: name.0, - email: email.0, - })) - } + self.db.create_user(id, name.0.clone(), email.0.clone(), hash.0).unwrap(); + Success(Json(User { + id, + username: name.0, + email: email.0, + })) } #[oai(path = "/user", method = "put")] /// Updates your current name and email - async fn update_user( - &self, - auth: Authorization, - name: Query<String>, - email: Query<String>, - ) -> CreateUserResponse { - use CreateUserResponse::*; - let id = auth.0.id; - let update_stmt = &stmt!(&format!( - "UPDATE {}.users SET name = '{}', email = '{}' WHERE id = {};", - self.kspc, name.0, email.0, id - )); - if let Err(e) = self.sess.execute(update_stmt).wait() { - InternalError(PlainText(e.to_string())) - } else { - Success(Json(User { - id, - username: name.0, - email: email.0, - })) - } + async fn update_user(&self, auth: Authorization, name: Query<String>, email: Query<String>) -> GenericResponse { + use GenericResponse::*; + self.db.update_user(auth.0.id, name.0, email.0).unwrap(); + Success } #[oai(path = "/user", method = "delete")] /// Deletes your user async fn delete_user(&self, auth: Authorization) -> DeleteResponse { use DeleteResponse::*; - let id = auth.0.id; - self.sess.execute(&stmt!(&format!( - "DELETE FROM {}.users WHERE id={};", self.kspc, id - ))).wait().unwrap(); - let res = self.sess.execute(&stmt!(&format!( - "SELECT groups FROM {}.user_groups WHERE id={};", - self.kspc, auth.0.id - ))).wait().unwrap(); - let groups: SetIterator = match res.row_count() { - 0 => return Success, - _ => res.first_row().unwrap().get(0).unwrap(), - }; - for group in groups { - self.__remove_group_member(group.get_i64().unwrap(), id); + self.db.delete_user(auth.0.id).unwrap(); + for group in self.db.get_user_groups(auth.0.id).unwrap() { + self.__remove_group_member(group, auth.0.id); } - self.sess.execute(&stmt!(&format!( - "DELETE FROM {}.user_groups WHERE id={};", self.kspc, id - ))).wait().unwrap(); + self.db.delete_user_groups(auth.0.id).unwrap(); Success } @@ -278,38 +160,19 @@ impl Api { /// Gets all groups accessible to you async fn get_groups(&self, auth: Authorization) -> GroupsResponse { use GroupsResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT groups FROM {}.user_groups WHERE id={};", - self.kspc, auth.0.id - ))).wait().unwrap(); - - let groups: SetIterator = match res.row_count() { - 0 => return Success(Json(Vec::new())), - _ => res.first_row().unwrap().get(0).unwrap(), - }; - - let group_vec = groups.map(|i| { - let res = self.sess.execute(&stmt!(&format!( - "SELECT id, name, members, channels FROM {}.groups WHERE id={};", self.kspc, i - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let (members, channels): (SetIterator, SetIterator) = (row.get(2).unwrap(), row.get(3).unwrap()); - Group { - id: row.get(0).unwrap(), - name: row.get(1).unwrap(), - members: members.map(|i| i.get_i64().unwrap()).collect(), - channels: channels.map(|i| i.get_i64().unwrap()).collect(), - } + let groups = self.db.get_user_groups(auth.0.id).unwrap(); + let group_vec = groups.iter().map(|i| { + self.db.get_group(*i).unwrap() }).collect(); - Success(Json(group_vec)) + Success(Json(group_vec)) } #[oai(path = "/user/groups", method = "delete")] /// Leaves a group accessible to you async fn leave_group(&self, auth: Authorization, gid: Query<i64>) -> GenericResponse { use GenericResponse::*; - if let Err(e) = self.validate_id("groups", gid.0) { - return NotFound(PlainText("Didn't find group or experienced database error.".to_string())); + if !self.db.valid_id(IdType::Group, gid.0).unwrap() { + return NotFound(PlainText("Group not found".to_string())); } self.__remove_group_member(gid.0, auth.0.id); Success @@ -319,24 +182,8 @@ impl Api { /// Gets the group with the given ID async fn get_group(&self, auth: Authorization, id: Query<i64>) -> GroupResponse { use GroupResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT name, members, channels FROM {}.groups WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - let (name, members, channels): (String, SetIterator, SetIterator) = match res.row_count() { - 1 => { - let row = res.first_row().unwrap(); - (row.get(0).unwrap(), row.get(1).unwrap(), row.get(2).unwrap()) - }, - 0 => return NotFound, - _ => return InternalError(PlainText(UNUSUAL_ROW_ERROR.to_string())) - }; - - Success(Json(Group{ - id: id.0, - name, - members: members.map(|i| i.get_i64().unwrap()).collect(), - channels: channels.map(|i| i.get_i64().unwrap()).collect(), - })) + if !self.db.valid_id(IdType::Group, id.0).unwrap() { return NotFound; } + Success(Json(self.db.get_group(id.0).unwrap())) } #[oai(path = "/group", method = "post")] @@ -347,19 +194,11 @@ impl Api { let cid = gen_id(); if name.0 == "" { return BadRequest(PlainText("Empty string not allowed for name".to_string())) - } - self.sess.execute(&stmt!(&format!( - "INSERT INTO {}.channels (id, group, name, members) VALUES ({}, {}, '{}', {{{}}});", - self.kspc, cid, gid, "main", auth.0.id - ))).wait().unwrap(); - self.sess.execute(&stmt!(&format!( - "INSERT INTO {}.groups (id, name, channels, members) VALUES ({}, '{}', {{{}}}, {{{}}});", - self.kspc, gid, name.0, cid, auth.0.id - ))).wait().unwrap(); - self.sess.execute(&stmt!(&format!( - "UPDATE {}.user_groups SET groups = groups + {{{}}} WHERE id = {};", - self.kspc, gid, auth.0.id - ))).wait().unwrap(); + } + self.db.create_group(gid, auth.0.id, name.0.clone()).unwrap(); + self.db.create_channel(cid, gid, auth.0.id, String::from("main")).unwrap(); + self.db.add_group_channel(gid, cid).unwrap(); + self.db.add_user_group(auth.0.id, gid).unwrap(); Success(Json(Group { id: gid, name: name.0, @@ -370,58 +209,32 @@ impl Api { #[oai(path = "/group", method = "put")] /// Updates the name of an existing group - async fn update_group( - &self, - auth: Authorization, - id: Query<i64>, - name: Query<String>, - ) -> GenericResponse { + async fn update_group(&self, auth: Authorization, id: Query<i64>, name: Query<String>) -> GenericResponse { use GenericResponse::*; if name.0 == "" { return BadRequest(PlainText("Empty string not allowed for name".to_string())) - } else if let Err(e) = self.validate_id("groups", id.0) { + } else if !self.db.valid_id(IdType::Group, id.0).unwrap() { return NotFound(PlainText("Didn't find group or experienced database error.".to_string())); } - self.sess.execute(&stmt!(&format!( - "UPDATE {}.groups SET name = '{}' WHERE id = {};", - self.kspc, name.0, id.0 - ))).wait().unwrap(); - Success + self.db.update_group(id.0, name.0).unwrap(); + Success } #[oai(path = "/group", method = "delete")] /// Deletes a group async fn delete_group(&self, auth: Authorization, id: Query<i64>) -> DeleteResponse { use DeleteResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT id, members, channels FROM {}.groups WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - - let (members, channels): (SetIterator, SetIterator) = match res.row_count() { - 1 => { - let row = res.first_row().unwrap(); - (row.get(1).unwrap(), row.get(2).unwrap()) - }, - 0 => return NotFound(PlainText("Group not found".to_string())), - _ => return InternalError(PlainText(UNUSUAL_ROW_ERROR.to_string())) - }; - - for member in members { - self.sess.execute(&stmt!(&format!( - "UPDATE {}.user_groups SET groups = groups - {{{}}} WHERE id = {};", - self.kspc, id.0, member - ))).wait().unwrap(); + if !self.db.valid_id(IdType::Group, id.0).unwrap() { + return NotFound(PlainText("Group not found".to_string())); } - - for channel in channels { - self.sess.execute(&stmt!(&format!( - "DELETE FROM {}.channels WHERE id={};", self.kspc, channel - ))).wait().unwrap(); + let group = self.db.get_group(id.0).unwrap(); + for member in group.members { + self.db.remove_user_group(member, id.0).unwrap(); + } + for channel in group.channels { + self.db.delete_channel(channel).unwrap(); } - - self.sess.execute(&stmt!(&format!( - "DELETE FROM {}.groups WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); + self.db.delete_group(id.0).unwrap(); Success } @@ -429,75 +242,36 @@ impl Api { /// Gets the members of the specified group async fn get_group_members(&self, auth: Authorization, id: Query<i64>) -> MembersResponse { use MembersResponse::*; - if let Err(_) = self.validate_id("groups", id.0) { + if !self.db.valid_id(IdType::Group, id.0).unwrap() { return NotFound; - } - let res = self.sess.execute(&stmt!(&format!( - "SELECT members FROM {}.groups WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let members: SetIterator = row.get(0).unwrap(); - let members_objs = members.map(|_| { - let res = self.sess.execute(&stmt!(&format!( - "SELECT id, name, email FROM {}.users WHERE id={};", - self.kspc, id.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - User { - id: id.0, - username: row.get(1).unwrap(), - email: row.get(2).unwrap(), - } - }).collect::<Vec<User>>(); - Success(Json(members_objs)) + } + let members = self.db.get_group_members(id.0).unwrap(); + Success(Json(members.iter().map(|m| { + self.db.get_user(*m).unwrap() + }).collect::<Vec<User>>())) } #[oai(path = "/group/members", method = "put")] /// Adds a member to an existing group - async fn add_group_member( - &self, - auth: Authorization, - gid: Query<i64>, - uid: Query<i64>, - ) -> GenericResponse { - use GenericResponse::*; - if let Err(e) = self.validate_id("groups", gid.0) { - return NotFound(PlainText("Didn't find group or experienced database error.".to_string())); + async fn add_group_member(&self, auth: Authorization, gid: Query<i64>, uid: Query<i64>) -> GenericResponse { + use GenericResponse::*; + if !self.db.valid_id(IdType::Group, gid.0).unwrap() { + return NotFound(PlainText("Group not found".to_string())); } - self.sess.execute(&stmt!(&format!( - "UPDATE {}.groups SET members = members + {{{}}} WHERE id = {};", - self.kspc, uid.0, gid.0 - ))).wait().unwrap(); - let res = self.sess.execute(&stmt!(&format!( - "SELECT channels FROM {}.groups WHERE id = {};", - self.kspc, gid.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let mut channels: SetIterator = row.get(0).unwrap(); - let cid: i64 = channels.next().unwrap().get_i64().unwrap(); - self.sess.execute(&stmt!(&format!( - "UPDATE {}.channels SET members = members + {{{}}} WHERE id = {};", - self.kspc, uid.0, cid - ))).wait().unwrap(); - self.sess.execute(&stmt!(&format!( - "UPDATE {}.user_groups SET groups = groups + {{{}}} WHERE id = {};", - self.kspc, gid.0, uid.0 - ))).wait().unwrap(); + self.db.add_group_member(gid.0, uid.0).unwrap(); + let channels = self.db.get_group_channels(gid.0).unwrap(); + self.db.add_channel_member(channels[0], uid.0).unwrap(); + self.db.add_user_group(uid.0, gid.0).unwrap(); Success } #[oai(path = "/group/members", method = "delete")] /// Removes a member from an existing group - async fn remove_group_member( - &self, - auth: Authorization, - gid: Query<i64>, - uid: Query<i64>, - ) -> DeleteResponse { + async fn remove_group_member(&self, auth: Authorization, gid: Query<i64>, uid: Query<i64>) -> DeleteResponse { use DeleteResponse::*; - if let Err(_) = self.validate_id("groups", gid.0) { + if !self.db.valid_id(IdType::Group, gid.0).unwrap() { return NotFound(PlainText("Group not found".to_string())) - } else if let Err(_) = self.validate_id("users", uid.0) { + } else if !self.db.valid_id(IdType::User, uid.0).unwrap() { return NotFound(PlainText("User not found".to_string())) } self.__remove_group_member(gid.0, uid.0); @@ -508,58 +282,27 @@ impl Api { /// Gets all channels in a group that are accessible to you async fn get_channels(&self, auth: Authorization, gid: Query<i64>) -> ChannelsResponse { use ChannelsResponse::*; - if let Err(_) = self.validate_id("groups", gid.0) { + if !self.db.valid_id(IdType::Group, gid.0).unwrap() { return NotFound; } - let res = self.sess.execute(&stmt!(&format!( - "SELECT channels FROM {}.groups WHERE id={};", self.kspc, gid.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let channels: SetIterator = row.get(0).unwrap(); - let mut channel_objs = Vec::new(); - for chan in channels { - let res = self.sess.execute(&stmt!(&format!( - "SELECT id, name, members FROM {}.channels WHERE id={};", self.kspc, chan - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let members: SetIterator = row.get(2).unwrap(); - let members = members.map(|m| m.get_i64().unwrap()).collect::<Vec<i64>>(); - if !members.contains(&auth.0.id) { - continue; - } - channel_objs.push(Channel { - id: row.get(0).unwrap(), - name: row.get(1).unwrap(), - group: gid.0, - members - }); - } - Success(Json(channel_objs)) + let channels = self.db.get_group_channels(gid.0).unwrap(); + Success(Json(channels.iter().map(|c| { + self.db.get_channel(*c).unwrap() + }).collect::<Vec<Channel>>())) } - + #[oai(path = "/group/channels", method = "post")] /// CREATES a channel in a group - async fn make_channel( - &self, - auth: Authorization, - gid: Query<i64>, - name: Query<String>, - ) -> CreateChannelResponse { + async fn make_channel(&self, auth: Authorization, gid: Query<i64>, name: Query<String>) -> CreateChannelResponse { use CreateChannelResponse::*; if name.0 == "" { return BadRequest(PlainText("Empty string not allowed for name".to_string())) - } else if let Err(e) = self.validate_id("groups", gid.0) { - return NotFound(PlainText("Group not found.".to_string())); - } + } else if !self.db.valid_id(IdType::Group, gid.0).unwrap() { + return NotFound(PlainText("Group not found".to_string())); + } let cid = gen_id(); - self.sess.execute(&stmt!(&format!( - "INSERT INTO {}.channels (id, group, name, members) VALUES ({}, {}, '{}', {{{}}});", - self.kspc, cid, gid.0, name.0, auth.0.id - ))).wait().unwrap(); - self.sess.execute(&stmt!(&format!( - "UPDATE {}.groups SET channels = channels + {{{}}} WHERE id = {};", - self.kspc, cid, gid.0 - ))).wait().unwrap(); + self.db.create_channel(cid, gid.0, auth.0.id, name.0.clone()).unwrap(); + self.db.add_group_channel(cid, gid.0).unwrap(); Success(Json(Channel { id: cid, name: name.0, @@ -570,25 +313,12 @@ impl Api { #[oai(path = "/channel", method = "put")] /// Updates the name of a channel - async fn update_channel( - &self, - auth: Authorization, - id: Query<i64>, - name: Query<String>, - ) -> GenericResponse { + async fn update_channel(&self, auth: Authorization, id: Query<i64>, name: Query<String>) -> GenericResponse { use GenericResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT id FROM {}.channels WHERE id = {};", self.kspc, id.0 - ))).wait().unwrap(); - match res.row_count() { - 0 => return NotFound(PlainText("Channel not found".to_string())), - 1 => (), - _ => return InternalError(PlainText(UNUSUAL_ROW_ERROR.to_string())) + if !self.db.valid_id(IdType::Channel, id.0).unwrap() { + return NotFound(PlainText("Channel not found".to_string())); } - self.sess.execute(&stmt!(&format!( - "UPDATE {}.channels SET name = '{}', WHERE id = {};", - self.kspc, name.0, id.0 - ))).wait().unwrap(); + self.db.update_channel(id.0, name.0).unwrap(); Success } @@ -596,169 +326,83 @@ impl Api { /// Gets a channel async fn get_channel(&self, auth: Authorization, id: Query<i64>) -> ChannelResponse { use ChannelResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT name, group, members FROM {}.channels WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - let (name, group, members): (String, i64, SetIterator) = match res.row_count() { - 1 => { - let row = res.first_row().unwrap(); - (row.get(0).unwrap(), row.get(1).unwrap(), row.get(2).unwrap()) - }, - 0 => return NotFound, - _ => return InternalError(PlainText(UNUSUAL_ROW_ERROR.to_string())) - }; - Success(Json(Channel { - id: id.0, - name, - group, - members: members.map(|i| i.get_i64().unwrap()).collect(), - })) + if !self.db.valid_id(IdType::Channel, id.0).unwrap() { + return NotFound; + } + Success(Json(self.db.get_channel(id.0).unwrap())) } - + #[oai(path = "/channel", method = "delete")] /// Deletes a channel async fn delete_channel(&self, auth: Authorization, id: Query<i64>) -> DeleteResponse { use DeleteResponse::*; - if let Err(_) = self.validate_id("channels", id.0) { - NotFound(PlainText("Channel not found.".to_string())) - } else { - let res = self.sess.execute(&stmt!(&format!( - "SELECT group FROM {}.channels WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let group: i64 = row.get(0).unwrap(); - self.sess.execute(&stmt!(&format!( - "UPDATE {}.groups SET channels = channels - {{{}}} WHERE id = {};", - self.kspc, id.0, group - ))).wait().unwrap(); - self.sess.execute(&stmt!(&format!( - "DELETE FROM {}.channels WHERE id = {};", - self.kspc, id.0 - ))).wait().unwrap(); - Success + if !self.db.valid_id(IdType::Channel, id.0).unwrap() { + return NotFound(PlainText("Channel not found".to_string())); } + let channel = self.db.get_channel(id.0).unwrap(); + self.db.remove_group_channel(channel.group, id.0).unwrap(); + self.db.delete_channel(id.0).unwrap(); + Success } #[oai(path = "/channel/members", method = "get")] /// Gets the members that can access a channel async fn get_channel_members(&self, auth: Authorization, id: Query<i64>) -> MembersResponse { use MembersResponse::*; - let res = self.sess.execute(&stmt!(&format!( - "SELECT members FROM {}.channels WHERE id={};", self.kspc, id.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - let members: SetIterator = row.get(0).unwrap(); - let members_objs = members.map(|_| { - let res = self.sess.execute(&stmt!(&format!( - "SELECT id, name, email FROM {}.users WHERE id={};", - self.kspc, id.0 - ))).wait().unwrap(); - let row = res.first_row().unwrap(); - User { - id: id.0, - username: row.get(1).unwrap(), - email: row.get(2).unwrap(), - } - }).collect::<Vec<User>>(); - Success(Json(members_objs)) + let members = self.db.get_channel_members(id.0).unwrap(); + Success(Json(members.iter().map(|m| { + self.db.get_user(*m).unwrap() + }).collect::<Vec<User>>())) } #[oai(path = "/channel/members", method = "put")] /// Adds a member to a channel - async fn add_channel_member( - &self, - auth: Authorization, - id: Query<i64>, - uid: Query<i64>, - ) -> GenericResponse { + async fn add_channel_member(&self, auth: Authorization, id: Query<i64>, uid: Query<i64>) -> GenericResponse { use GenericResponse::*; - if let Err(_) = self.validate_id("channels", id.0) { + if !self.db.valid_id(IdType::Channel, id.0).unwrap() { return NotFound(PlainText("Channel not found".to_string())) - } else if let Err(_) = self.validate_id("users", uid.0) { + } else if !self.db.valid_id(IdType::User, uid.0).unwrap() { return NotFound(PlainText("User not found".to_string())) } - self.sess.execute(&stmt!(&format!( - "UPDATE {}.channels SET members = members + {{{}}} WHERE id={};", self.kspc, uid.0, id.0 - ))).wait().unwrap(); + self.db.add_channel_member(id.0, uid.0).unwrap(); Success } #[oai(path = "/channel/members", method = "delete")] /// Removes a member from a channel - async fn remove_channel_member( - &self, - auth: Authorization, - cid: Query<i64>, - uid: Query<i64>, - ) -> DeleteResponse { + async fn remove_channel_member(&self, auth: Authorization, cid: Query<i64>, uid: Query<i64>) -> DeleteResponse { use DeleteResponse::*; - if let Err(_) = self.validate_id("channels", cid.0) { + if !self.db.valid_id(IdType::Channel, cid.0).unwrap() { return NotFound(PlainText("Channel not found".to_string())) - } else if let Err(_) = self.validate_id("users", uid.0) { + } else if !self.db.valid_id(IdType::User, uid.0).unwrap() { return NotFound(PlainText("User not found".to_string())) } - self.__remove_channel_member(cid.0, uid.0); + self.db.remove_channel_member(cid.0, uid.0).unwrap(); Success } - #[oai(path = "/channel/message", method = "get")] + #[oai(path = "/channel/term", method = "get")] /// Returns batch of messages in channel containing "term" in the last 100 messages - async fn search_channel( - &self, - auth: Authorization, - cid: Query<i64>, - term: Query<String>, - off: Query<u64>, - ) -> MessagesResponse { + async fn search_channel(&self, auth: Authorization, cid: Query<i64>, term: Query<String>, off: Query<u64>) -> MessagesResponse { use MessagesResponse::*; - if let Err(_) = self.validate_id("channels", cid.0) { + if !self.db.valid_id(IdType::Channel, cid.0).unwrap() { return NotFound(PlainText("Channel not found".to_string())) } - let res = self.sess.execute(&stmt!(&format!( - "SELECT * FROM {}.messages WHERE channel={} LIMIT 100;", - self.kspc, cid.0 - ))).wait().unwrap(); - let messages = res.iter().filter(|row| { - let content: String = row.get(4).unwrap(); - content.contains(&term.0) - }).map(|row| { - Message { - id: row.get(0).unwrap(), - author: row.get(2).unwrap(), - channel: row.get(1).unwrap(), - content: row.get(4).unwrap() - } - }).collect::<Vec<Message>>(); - Success(Json(messages)) + let mut messages = self.db.get_messages(cid.0, 100).unwrap(); + messages.retain(|msg| msg.content.contains(&term.0)); + Success(Json(messages)) } #[oai(path = "/channel/messages", method = "get")] /// Returns batch of messages in channel. Do not use for small batches. /// /// For small batches, use `chatterbox`, the websocket service for messaging, instead. - async fn get_channel_messages( - &self, - auth: Authorization, - cid: Query<i64>, - num_msgs: Query<u64>, - ) -> MessagesResponse { + async fn get_channel_messages(&self, auth: Authorization, cid: Query<i64>, num_msgs: Query<u64>) -> MessagesResponse { use MessagesResponse::*; - if let Err(_) = self.validate_id("channels", cid.0) { + if !self.db.valid_id(IdType::Channel, cid.0).unwrap() { return NotFound(PlainText("Channel not found".to_string())) } - let res = self.sess.execute(&stmt!(&format!( - "SELECT * FROM {}.messages WHERE channel={} LIMIT {};", - self.kspc, cid.0, num_msgs.0, - ))).wait().unwrap(); - let messages = res.iter().map(|row| { - Message { - id: row.get(0).unwrap(), - author: row.get(2).unwrap(), - channel: row.get(1).unwrap(), - content: row.get(4).unwrap() - } - }).collect::<Vec<Message>>(); - Success(Json(messages)) + Success(Json(self.db.get_messages(cid.0, num_msgs.0).unwrap())) } } @@ -770,7 +414,8 @@ async fn main() -> Result<(), std::io::Error> { } tracing_subscriber::fmt::init(); - let api_service = OpenApiService::new(Api::new("bsk"), "Scuttlebutt", "1.0") + let db = Box::new(Cassandra::new("bsk")); + let api_service = OpenApiService::new(Api::new(db), "Scuttlebutt", "1.0") .description( "Scuttlebutt is the REST API for managing everything but sending/receiving messages \ - which means creating/updating/deleting all of your users/groups/channels.", @@ -778,7 +423,6 @@ async fn main() -> Result<(), std::io::Error> { .server("http://localhost:3000/api"); let ui = api_service.swagger_ui(); - let spec = api_service.spec(); let key: String = rand::thread_rng() .sample_iter(&Alphanumeric) @@ -790,15 +434,8 @@ async fn main() -> Result<(), std::io::Error> { .nest("/api", api_service) .nest("/", ui) .data(ServerKey::new_from_slice(&key.as_bytes()).unwrap()); - // let cli = poem::test::TestClient::new(app); - // let resp = cli.post("/api/login?id=234").body("abc").send(); - // resp.assert_status_is_ok(); - // Ok(()) - Server::new(TcpListener::bind("127.0.0.1:3000")) - .run( - app, // .at("/spec", poem::endpoint::make_sync(move |_| spec.clone())) - ) - .await + + Server::new(TcpListener::bind("127.0.0.1:3000")).run(app).await } #[cfg(test)] diff --git a/scuttlebutt/src/tests.rs b/scuttlebutt/src/tests.rs @@ -5,7 +5,7 @@ use more_asserts::*; use poem::{ http::StatusCode, middleware::AddDataEndpoint, - test::{TestClient, TestRequestBuilder, TestResponse}, + test::TestClient, Route, }; use pretty_assertions::assert_eq; @@ -13,13 +13,18 @@ use sha2::Digest; type FakeClient = TestClient<AddDataEndpoint<Route, ServerKey>>; +fn contents_eq<T: PartialEq>(a: Vec<T>, b: Vec<T>) -> bool { + b.iter().all(|item| a.contains(item)) +} + fn setup() -> FakeClient { let key: String = rand::thread_rng() .sample_iter(&Alphanumeric) .take(7) .map(char::from) .collect(); - let api_service = OpenApiService::new(Api::new("test"), "Scuttlebutt", "1.0").server("http://localhost:3000/api"); + let db = Box::new(Cassandra::new("test")); + let api_service = OpenApiService::new(Api::new(db), "Scuttlebutt", "1.0").server("http://localhost:3000/api"); let app = Route::new() .nest("/api", api_service) .data(ServerKey::new_from_slice(&key.as_bytes()).unwrap()); @@ -32,14 +37,6 @@ fn hash_pass(pass: &str) -> String { hex::encode(hasher.finalize()) } -enum HttpMethod { - Post(String), - Get(String), - Delete(String), - Put(String), -} -use crate::tests::HttpMethod::*; - async fn make_user(cli: &FakeClient, name: &str, email: &str, pass: &str) -> User { let hash = hash_pass(pass); let resp = cli @@ -70,16 +67,24 @@ async fn setup_user_auth() -> (FakeClient, User, String) { /// FIXME non exhaustive async fn post_user() { let cli = setup(); - let resp = make_user(&cli, "test", "test@example.com", "12345").await; + let user = make_user(&cli, "test", "test@example.com", "12345").await; + + assert_eq!(user.email, "test@example.com"); + assert_eq!(user.username, "test"); // TODO questionable - let mut id_gen = Snowflake::default(); - assert_ge!(id_gen.generate(), resp.id); + // let mut id_gen = Snowflake::default(); + // assert_ge!(id_gen.generate(), resp.id); - assert_eq!(resp.email, "test@example.com"); - assert_eq!(resp.username, "test"); + let resp = cli.get(format!("/api/user?id={}", user.id)).send().await; + resp.assert_status_is_ok(); + + let same_user = resp.json().await.value().deserialize::<User>(); + assert_eq!(user, same_user); } + + #[tokio::test] async fn post_login() { let cli = setup(); @@ -122,15 +127,16 @@ async fn get_user() { #[tokio::test] async fn put_user() { let (cli, user, auth) = setup_user_auth().await; - + let resp = cli.put("/api/user?name=fred&email=whoo@whee.com").send().await; resp.assert_status(StatusCode::UNAUTHORIZED); - let resp = cli - .put("/api/user?name=fred&email=whoo@whee.com") - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; + + let resp = cli.put("/api/user?name=fred&email=whoo@whee.com") + .header::<&str, String>("ScuttleKey", auth).send().await; resp.assert_status_is_ok(); + + let resp = cli.get(format!("/api/user?id={}", user.id)).send().await; + resp.assert_status_is_ok(); let ret_user = resp.json().await.value().deserialize::<User>(); assert_eq!(ret_user.email, "whoo@whee.com"); assert_eq!(ret_user.username, "fred"); @@ -145,50 +151,45 @@ async fn del_user() { let resp = cli.delete(format!("/api/user?id={}", user.id)).send().await; resp.assert_status(StatusCode::UNAUTHORIZED); - let resp = cli - .delete(format!("/api/user?id={}", user.id)) - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; + let resp = cli.delete(format!("/api/user?id={}", user.id)) + .header::<&str, String>("ScuttleKey", auth).send().await; + resp.assert_status_is_ok(); let resp = cli.get(format!("/api/user?id={}", user.id)).send().await; resp.assert_status(StatusCode::NOT_FOUND); } -async fn make_group(cli: &FakeClient, auth: String, name: &str) -> Group { - let resp = cli - .post(format!("/api/group?name={}", name)) - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; +async fn make_group(cli: &FakeClient, auth: &str, name: &str) -> Group { + let resp = cli.post(format!("/api/group?name={}", name)) + .header::<&str, &str>("ScuttleKey", auth).send().await; resp.assert_status_is_ok(); resp.json().await.value().deserialize::<Group>() } -async fn make_channel(cli: &FakeClient, auth: String, gid: i64, name: &str) -> Channel { +async fn make_channel(cli: &FakeClient, auth: &str, gid: i64, name: &str) -> Channel { let resp = cli .post(format!("/api/group/channels?gid={}&name={}", gid, name)) - .header::<&str, String>("ScuttleKey", auth) + .header::<&str, &str>("ScuttleKey", auth) .send() .await; resp.assert_status_is_ok(); resp.json().await.value().deserialize::<Channel>() } -async fn find_channel(cli: &FakeClient, auth: String, id: i64) -> Channel { +async fn find_channel(cli: &FakeClient, auth: &str, id: i64) -> Channel { let resp = cli .get(format!("/api/channel?id={}", id)) - .header::<&str, String>("ScuttleKey", auth) + .header::<&str, &str>("ScuttleKey", auth) .send() .await; resp.assert_status_is_ok(); resp.json().await.value().deserialize::<Channel>() } -async fn find_group(cli: &FakeClient, auth: String, id: i64) -> Group { +async fn find_group(cli: &FakeClient, auth: &str, id: i64) -> Group { let resp = cli .get(format!("/api/group?id={}", id)) - .header::<&str, String>("ScuttleKey", auth) + .header::<&str, &str>("ScuttleKey", auth) .send() .await; resp.assert_status_is_ok(); @@ -199,12 +200,12 @@ async fn find_group(cli: &FakeClient, auth: String, id: i64) -> Group { async fn post_group() { let (cli, user, auth) = setup_user_auth().await; let resp = cli.post("/api/group?name=") - .header::<&str, String>("ScuttleKey", auth.clone()).send().await; + .header::<&str, &str>("ScuttleKey", &auth).send().await; resp.assert_status(StatusCode::BAD_REQUEST); let resp = cli .post("/api/group?name=test") - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status_is_ok(); @@ -213,7 +214,7 @@ async fn post_group() { assert_eq!(group.members, vec![user.id]); assert_eq!(group.channels.len(), 1); assert_eq!( - find_channel(&cli, auth, group.channels[0]).await.name, + find_channel(&cli, &auth, group.channels[0]).await.name, String::from("main") ); } @@ -221,67 +222,64 @@ async fn post_group() { #[tokio::test] async fn put_group() { let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; + let group = make_group(&cli, &auth, "test").await; let resp = cli.put(format!("/api/group?id={}&name=test2", group.id)).send().await; resp.assert_status(StatusCode::UNAUTHORIZED); let resp = cli .put(format!("/api/group?id={}&name=", group.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status(StatusCode::BAD_REQUEST); let resp = cli .put("/api/group?id=12&name=test2") - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status(StatusCode::NOT_FOUND); - let resp = cli - .put(format!("/api/group?id={}&name=test2", group.id)) - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; + let resp = cli.put(format!("/api/group?id={}&name=test2", group.id)) + .header::<&str, String>("ScuttleKey", auth).send().await; resp.assert_status_is_ok(); } #[tokio::test] async fn del_group() { let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; + let group = make_group(&cli, &auth, "test").await; let resp = cli.delete(format!("/api/group?id={}", group.id)).send().await; resp.assert_status(StatusCode::UNAUTHORIZED); - let resp = cli - .delete("/api/group?id=12") - .header::<&str, String>("ScuttleKey", auth.clone()) - .send() - .await; + let resp = cli.delete("/api/group?id=12") + .header::<&str, &str>("ScuttleKey", &auth).send().await; resp.assert_status(StatusCode::NOT_FOUND); - let resp = cli - .delete(format!("/api/group?id={}", group.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) - .send() - .await; + let resp = cli.delete(format!("/api/group?id={}", group.id)) + .header::<&str, &str>("ScuttleKey", &auth).send().await; resp.assert_status_is_ok(); - let resp = cli - .get(format!("/api/group?id={}", group.id)) - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; + let resp = cli.get(format!("/api/group?id={}", group.id)) + .header::<&str, &str>("ScuttleKey", &auth).send().await; + + resp.assert_status(StatusCode::NOT_FOUND); + let resp = cli.get(format!("/api/channel?id={}", group.channels[0])) + .header::<&str, &str>("ScuttleKey", &auth).send().await; + resp.assert_status(StatusCode::NOT_FOUND); + + let resp = cli.get("/api/user/groups") + .header::<&str, String>("ScuttleKey", auth).send().await; + resp.assert_status(StatusCode::NOT_FOUND); } #[tokio::test] async fn put_group_members() { let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; + let group = make_group(&cli, &auth, "test").await; let user2 = make_user(&cli, "testeroo", "test2@example.com", "123456").await; let user3 = make_user(&cli, "testeroo", "test2@example.com", "123456").await; @@ -293,34 +291,34 @@ async fn put_group_members() { let resp = cli .put("/api/group/members?gid=12&uid=32") - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status(StatusCode::NOT_FOUND); let resp = cli .put(format!("/api/group/members?gid={}&uid={}", group.id, user2.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status_is_ok(); let resp = cli .put(format!("/api/group/members?gid={}&uid={}", group.id, user3.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status_is_ok(); - let new_group = find_group(&cli, auth, group.id).await; - assert_eq!(new_group.members, vec![user.id, user2.id, user3.id]); + let new_group = find_group(&cli, &auth, group.id).await; + assert!(contents_eq(new_group.members, vec![user.id, user2.id, user3.id])); } #[tokio::test] async fn post_group_channels() { let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; + let group = make_group(&cli, &auth, "test").await; - let resp = cli + let resp = cli .post(format!("/api/group/channels?gid={}&name=test", group.id)) .send() .await; @@ -328,34 +326,35 @@ async fn post_group_channels() { let resp = cli .post(format!("/api/group/channels?gid={}&name=", group.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status(StatusCode::BAD_REQUEST); let resp = cli .post("/api/group/channels?gid=12&name=test") - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status(StatusCode::NOT_FOUND); let resp = cli .post(format!("/api/group/channels?gid={}&name=test", group.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) + .header::<&str, &str>("ScuttleKey", &auth) .send() .await; resp.assert_status_is_ok(); let channel = resp.json().await.value().deserialize::<Channel>(); assert_eq!(channel.name, "test"); assert_eq!(channel.members, vec![user.id]); - assert!(find_group(&cli, auth, group.id).await.channels.contains(&channel.id)); + assert!(find_group(&cli, &auth, group.id).await.channels.contains(&channel.id)); } #[test] /// Test if gen_id() gives unique IDs on successive calls /// and if it can be called from multiple threads without error fn test_id_gen() { let a = gen_id(); + std::thread::sleep(std::time::Duration::from_secs(1)); let b = gen_id(); assert_ge!(b, a); let threads: Vec<_> = (0..100).map(|i| std::thread::spawn(move || gen_id())).collect(); @@ -366,14 +365,11 @@ fn test_id_gen() { #[tokio::test] async fn get_channel() { - let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; - let chan = make_channel(&cli, auth.clone(), group.id, "random").await; - let resp = cli - .get(format!("/api/channel?id={}", chan.id)) - .header::<&str, String>("ScuttleKey", auth) - .send() - .await; + let (cli, _user, auth) = setup_user_auth().await; + let group = make_group(&cli, &auth, "test").await; + let chan = make_channel(&cli, &auth, group.id, "random").await; + let resp = cli.get(format!("/api/channel?id={}", chan.id)) + .header::<&str, &str>("ScuttleKey", &auth).send().await; resp.assert_status_is_ok(); let recv_chan = resp.json().await.value().deserialize::<Channel>(); assert_eq!(chan, recv_chan); @@ -382,22 +378,32 @@ async fn get_channel() { // FIXME non exhaustive #[tokio::test] async fn get_channels() { - let (cli, user, auth) = setup_user_auth().await; - let group = make_group(&cli, auth.clone(), "test").await; - let chan1 = make_channel(&cli, auth.clone(), group.id, "random").await; - let chan2 = make_channel(&cli, auth.clone(), group.id, "random").await; - let chan3 = make_channel(&cli, auth.clone(), group.id, "random").await; - let resp = cli - .get(format!("/api/group/channels?gid={}", group.id)) - .header::<&str, String>("ScuttleKey", auth.clone()) - .send() - .await; + let (cli, _user, auth) = setup_user_auth().await; + let group = make_group(&cli, &auth, "test").await; + let chan1 = make_channel(&cli, &auth, group.id, "random").await; + let chan2 = make_channel(&cli, &auth, group.id, "random").await; + let chan3 = make_channel(&cli, &auth, group.id, "random").await; + let resp = cli.get(format!("/api/group/channels?gid={}", group.id)) + .header::<&str, &str>("ScuttleKey", &auth).send().await; resp.assert_status_is_ok(); let channels = resp.json().await.value().deserialize::<Vec<Channel>>(); - assert_eq!( + assert!(contents_eq( channels, - vec![find_channel(&cli, auth, group.channels[0]).await, chan1, chan2, chan3] - ); + vec![find_channel(&cli, &auth, group.channels[0]).await, chan1, chan2, chan3] + )); +} + +#[tokio::test] +async fn get_groups() { + let (cli, _user, auth) = setup_user_auth().await; + let group = make_group(&cli, &auth, "test1").await; + let group2 = make_group(&cli, &auth, "test2").await; + let group3 = make_group(&cli, &auth, "test3").await; + let resp = cli.get("/api/user/groups") + .header::<&str, &str>("ScuttleKey", &auth).send().await; + resp.assert_status_is_ok(); + let groups = resp.json().await.value().deserialize::<Vec<Group>>(); + assert!(contents_eq(groups, vec![group, group2, group3])); } // #[tokio::test]