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:
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]