blatherskite

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

commit 77941234e41781b75f21693b4b012ad0c1693d6c
parent 422c4fe5ae18daa4866e3033bf70034555d49dfb
Author: quantumish <freifeld.david@gmail.com>
Date:   Sun, 25 Sep 2022 22:09:14 -0700

Add barebones databasing, tests, id generation

Diffstat:
Mscuttlebutt/Cargo.toml | 6+++++-
Mscuttlebutt/src/main.rs | 225+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++------------
Mscuttlebutt/src/responses.rs | 22++++++++++++++--------
3 files changed, 212 insertions(+), 41 deletions(-)

diff --git a/scuttlebutt/Cargo.toml b/scuttlebutt/Cargo.toml @@ -7,13 +7,17 @@ edition = "2021" [dependencies] base64 = "0.13.0" +cassandra-cpp = "1.1.0" +chrono = { version = "0.4.22", features = ["serde"] } ctor = "0.1.23" hmac = "0.12.1" jwt = "0.16.0" log = "0.4.17" +more-asserts = "0.3.0" poem = { version = "1.3.42", features = ["test"] } poem-openapi = { version = "2.0.12", features = ["swagger-ui"] } -rs-snowflake = "0.6.0" +pretty_assertions = "1.3.0" +rustflake = "0.1.1" serde = "1.0.144" serde_json = "1.0.85" sha2 = "0.10.6" diff --git a/scuttlebutt/src/main.rs b/scuttlebutt/src/main.rs @@ -1,6 +1,8 @@ +use cassandra_cpp::*; +use chrono::{DateTime, Duration, Local}; use hmac::{Hmac, Mac}; use jwt::{SignWithKey, VerifyWithKey}; -use poem::{listener::TcpListener, Request, Result, Route, Server}; +use poem::{listener::TcpListener, web::Data, EndpointExt, Request, Result, Route, Server}; use poem_openapi::{ auth::ApiKey, param::Query, @@ -9,12 +11,20 @@ use poem_openapi::{ }; use serde::{Deserialize, Serialize}; use sha2::Sha256; +use std::sync::{Arc, Mutex}; pub mod responses; pub use responses::*; +use rustflake::Snowflake; type ServerKey = Hmac<Sha256>; +#[derive(Serialize, Deserialize)] +struct Claims { + id: i64, + exp: DateTime<Local>, +} + /// ApiKey authorization #[derive(SecurityScheme)] #[oai( @@ -26,28 +36,58 @@ type ServerKey = Hmac<Sha256>; struct Authorization(User); async fn api_checker(req: &Request, api_key: ApiKey) -> Option<User> { - let claims: User = serde_json::from_str( + let claims: Claims = serde_json::from_str( &String::from_utf8(base64::decode(api_key.key.split(".").nth(1).unwrap()).unwrap()) .unwrap(), ) .unwrap(); let key = todo!(); // Query DB here... - let server_key = req.data::<ServerKey>().unwrap(); - VerifyWithKey::<User>::verify_with_key(api_key.key.as_str(), server_key).ok() + // let server_key = req.data::<ServerKey>().unwrap(); + // VerifyWithKey::<User>::verify_with_key(api_key.key.as_str(), server_key).ok() +} + +struct Api { + sess: Session, + kspc: String, } -struct Api; +fn gen_id() -> i64 { + static STATE: Mutex<Option<Snowflake>> = Mutex::new(None); + + STATE + .lock() + .unwrap() + .get_or_insert_with(|| Snowflake::default()) + .generate() +} #[OpenApi] 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(); + + let space_stmt = &stmt!(&format!("CREATE KEYSPACE IF NOT EXISTS {keyspc} WITH replication = {{'class':'SimpleStrategy', 'replication_factor': 1}}")); + let table_stmt = &stmt!(&format!("CREATE TABLE IF NOT EXISTS {keyspc}.users (id bigint PRIMARY KEY, name text, email text, hash text);")); + session.execute(space_stmt).wait().unwrap(); + session.execute(table_stmt).wait().unwrap(); + + Api { + sess: session, + kspc: String::from(keyspc), + } + } + #[oai(path = "/login", method = "post")] - async fn login(&self, id: Query<String>, hash: Base64<Vec<u8>>) -> Result<PlainText<String>> { + async fn login(&self, id: Query<i64>, hash: Base64<Vec<u8>>) -> Result<PlainText<String>> { let key = Hmac::<Sha256>::new_from_slice(&hash.0).map_err(poem::error::InternalServerError)?; - let token = User { - username: String::from("Blurgh"), - email: String::from("Blurgh"), - id: 0, + let token = Claims { + id: id.0, + exp: Local::now() + Duration::days(1), } .sign_with_key(&key) .map_err(poem::error::InternalServerError)?; @@ -61,7 +101,30 @@ impl Api { /// /// Call `/user?id=1234` to get the user with id 1234 async fn get_user(&self, id: Query<i64>) -> UserResponse { - todo!() + 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( + "Found duplicate ID (or negative rows??)! Giving up!".to_string(), + )) + } + } + Err(e) => InternalError(PlainText(e.to_string())), + } } #[oai(path = "/user", method = "post")] @@ -72,7 +135,24 @@ impl Api { email: Query<String>, hash: Query<String>, ) -> CreateUserResponse { - todo!() + 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, + })) + } } #[oai(path = "/user", method = "put")] @@ -256,7 +336,7 @@ async fn main() -> Result<(), std::io::Error> { } tracing_subscriber::fmt::init(); - let api_service = OpenApiService::new(Api, "Scuttlebutt", "1.0") + let api_service = OpenApiService::new(Api::new("bsk"), "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.", @@ -264,31 +344,112 @@ async fn main() -> Result<(), std::io::Error> { .server("http://localhost:3000/api"); let ui = api_service.swagger_ui(); + let spec = api_service.spec(); - // let wat = b"whee"; - // let server_key = Hmac::<Sha256>::new_from_slice(wat).expect("valid server key"); - // let server_key2 = Hmac::<Sha256>::new_from_hex - // println!("{:?}", server_key.); Server::new(TcpListener::bind("127.0.0.1:3000")) - .run(Route::new().nest("/api", api_service).nest("/", ui)) + .run( + Route::new() + .nest("/api", api_service) + .nest("/", ui) + .at("/spec", poem::endpoint::make_sync(move |_| spec.clone())), + ) .await } #[cfg(test)] mod tests { use super::*; - use poem::test::TestClient; - - // fn setup() -> TestClient<Route> { - // let app = OpenApiService::new(Api, "Scuttlebutt", "1.0").server("http://localhost:3000/api"); - // TestClient::new(Route::new().nest("/api", app)) - // } - - // #[tokio::test] - // async fn sanity() { - // let cli = setup(); - // let resp = cli.get("/api/hello").send().await; - // resp.assert_status_is_ok(); - // resp.assert_text("whee").await; - // } + use more_asserts::*; + use poem::{ + http::StatusCode, + test::{TestClient, TestResponse}, + }; + use pretty_assertions::{assert_eq, assert_ne}; + use sha2::Digest; + + fn setup() -> TestClient<Route> { + let app = + OpenApiService::new(Api, "Scuttlebutt", "1.0").server("http://localhost:3000/api"); + TestClient::new(Route::new().nest("/api", app)) + } + + fn hash_pass(pass: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(pass); + base64::encode(hasher.finalize()) + } + + async fn make_user(cli: &TestClient<Route>, name: &str, email: &str, pass: &str) -> User { + let mut hasher = Sha256::new(); + hasher.update(pass); + let hash = hash_pass(pass); + let resp = cli + .post(format!( + "/api/user?name={}?email={}?hash={}", + name, email, hash + )) + .send() + .await; + resp.assert_status_is_ok(); + resp.json().await.value().deserialize::<User>() + } + + #[tokio::test] + async fn test_make_user() { + let cli = setup(); + let resp = make_user(&cli, "test", "test@example.com", "12345").await; + + // TODO questionable + let mut id_gen = snowflake::SnowflakeIdGenerator::new(1, 1); + assert_ge!(id_gen.real_time_generate(), resp.id); + + assert_eq!(resp.email, "test@example.com"); + assert_eq!(resp.username, "test"); + } + + #[tokio::test] + async fn test_login_bad_hash() { + let cli = setup(); + let resp = cli + .post("/api/user?name=test?email=test@example.com?hash=abc") + .send() + .await; + resp.assert_status(StatusCode::BAD_REQUEST); + } + + #[tokio::test] + async fn test_get_user() { + let cli = setup(); + let user = make_user(&cli, "test", "test@example.com", "12345").await; + 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!(user, ret_user); + } + + #[tokio::test] + async fn test_put_user() { + let cli = setup(); + let user = make_user(&cli, "test", "test@example.com", "12345").await; + let resp = cli + .put(format!("/api/user?name=fred?email=whoo@whee.com?hash=abc")) + .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"); + // User should retain their underlying ID + assert_eq!(user.id, ret_user.id); + } + + #[tokio::test] + async fn test_del_user() { + let cli = setup(); + let user = make_user(&cli, "test", "test@example.com", "12345").await; + let resp = cli.delete(format!("/api/user?id={}", user.id)).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); + } } diff --git a/scuttlebutt/src/responses.rs b/scuttlebutt/src/responses.rs @@ -4,14 +4,14 @@ use poem_openapi::{ }; use serde::{Deserialize, Serialize}; -#[derive(Object, Serialize, Deserialize)] +#[derive(Object, Serialize, Deserialize, Debug, Eq, PartialEq)] pub struct User { pub id: i64, pub username: String, pub email: String, } -#[derive(Object)] +#[derive(Object, Serialize, Deserialize, Debug, Eq, PartialEq)] pub struct Group { pub id: i64, pub name: String, @@ -21,14 +21,14 @@ pub struct Group { pub channels: Vec<i64>, } -#[derive(Object)] +#[derive(Object, Serialize, Deserialize, Debug, Eq, PartialEq)] pub struct Channel { pub id: i64, pub name: String, pub members: Vec<i64>, } -#[derive(Object)] +#[derive(Object, Serialize, Deserialize, Debug, Eq, PartialEq)] pub struct Message { pub id: i64, pub channel: i64, @@ -40,22 +40,28 @@ pub struct Message { pub enum UserResponse { /// Returns the user requested. #[oai(status = 200)] - User(Json<User>), + Success(Json<User>), /// Invalid ID. Content specifies which of the IDs passed is invalid. #[oai(status = 404)] - NotFound(PlainText<String>), + NotFound, + /// Internal server error: likely due to a database operation failing + #[oai(status = 500)] + InternalError(PlainText<String>), } #[derive(ApiResponse)] pub enum CreateUserResponse { /// Returns the user requested. #[oai(status = 200)] - User(Json<User>), + Success(Json<User>), /// Recieved a bad argument when specifying the user. Returns error type, such as: /// - found empty string for any of the arguments /// - invalid email #[oai(status = 400)] - BadRequest, + BadRequest(PlainText<String>), + /// Internal server error: likely due to a database operation failing + #[oai(status = 500)] + InternalError(PlainText<String>), } #[derive(ApiResponse)]