blatherskite

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

main.rs (7777B)


      1 /// Currently a modified version of `poem`'s default websocket-chat example
      2 use cassandra_cpp::*;
      3 use futures_util::{SinkExt, StreamExt};
      4 use poem::{
      5     get, handler,
      6     listener::TcpListener,
      7     web::{
      8         websocket::{Message, WebSocket},
      9         Data, Path,
     10     },
     11     EndpointExt, IntoResponse, Route, Server,
     12 };
     13 use rustflake::Snowflake;
     14 use serde_json::Value;
     15 use std::result::Result;
     16 use serde::{Serialize, Deserialize};
     17 use chrono::{DateTime, Local};
     18 
     19 #[derive(Serialize, Deserialize, Debug)]
     20 pub struct MessageObj {
     21     pub id: i64,
     22     pub channel: i64,
     23     pub author: i64,
     24     pub content: String,
     25 }
     26 
     27 pub fn gen_id() -> i64 {
     28     static STATE: std::sync::Mutex<Option<Snowflake>> = std::sync::Mutex::new(None);
     29 
     30     STATE
     31         .lock()
     32         .unwrap()
     33         .get_or_insert_with(|| Snowflake::new(1_564_790_400_000, 2, 1))
     34         .generate()
     35 }
     36 
     37 const KEYSPC: &'static str = "bsk";
     38 
     39 fn setup_db() -> Session {
     40     let contact_points = "127.0.0.1";
     41     let mut cluster = Cluster::default();
     42     cluster.set_contact_points(contact_points).unwrap();
     43     cluster.set_load_balance_round_robin();
     44     cluster.connect().unwrap()
     45 }
     46 
     47 #[handler]
     48 fn ws(  
     49     ws: WebSocket,
     50     sender: Data<&tokio::sync::broadcast::Sender<String>>,
     51 ) -> impl IntoResponse {
     52     let sender = sender.clone();
     53     ws.on_upgrade(move |socket| async move {
     54     let mut receiver = sender.subscribe();
     55         let (mut sink, mut stream) = socket.split();
     56 
     57         tokio::spawn(async move {
     58             let sess = setup_db();
     59             let mut user: Option<Value> = None;
     60             while let Some(Ok(msg)) = stream.next().await {
     61                 if let Message::Text(auth) = msg {                  
     62                     let req: Value = serde_json::from_str(&auth).unwrap();
     63                     let res = sess.execute(&stmt!(&format!(
     64                         "SELECT hash FROM {}.users WHERE id={};",
     65                         KEYSPC, req["id"].as_i64().unwrap(),
     66                     ))).wait().unwrap();
     67                     let row = res.first_row().unwrap();
     68                     let db_hash: String = row.get(0).unwrap();
     69                     if hex::decode(db_hash).unwrap() != hex::decode(req["hash"].as_str().unwrap()).unwrap() {
     70                         return;
     71                     }
     72                     user = Some(req);
     73                     break
     74                 }
     75             }
     76             while let Some(Ok(mesg)) = stream.next().await {
     77                 if let Message::Text(text) = mesg {
     78                     let id = gen_id();
     79                     let req: Value = serde_json::from_str(&text).unwrap();
     80                     let msg = MessageObj {
     81                         id,
     82                         content: req["content"].as_str().unwrap().to_string(),
     83                         author: user.clone().unwrap()["id"].as_i64().unwrap(),
     84                         channel: req["channel"].as_i64().unwrap(),
     85                     };
     86                     let mut stmt = stmt!(&format!(
     87                         "INSERT INTO {}.messages (channel, id, author, content) VALUES ({},{},{},?);",
     88                         KEYSPC, msg.channel, gen_id(), msg.author
     89                     ));
     90 					stmt.bind(0, msg.content.as_str()).unwrap();
     91 				    sess.execute(&stmt).wait().unwrap();
     92                     if sender.send(serde_json::to_string(&msg).unwrap()).is_err() {
     93                         break;
     94                     }
     95                 }
     96             }
     97         });
     98 
     99         tokio::spawn(async move {
    100             let sess = setup_db();
    101             while let Ok(msg) = receiver.recv().await {
    102                 let req: Value = serde_json::from_str(&msg).unwrap();
    103                 let res = sess.execute(&stmt!(&format!(
    104                     "SELECT members FROM {}.channels WHERE id={};", KEYSPC, req["channel"].as_i64().unwrap(),
    105                 ))).wait().unwrap();
    106                 let row = res.first_row().unwrap();
    107                 let members: SetIterator = row.get(0).unwrap();
    108                 if !members.map(|i| i.get_i64().unwrap()).collect::<Vec<i64>>().contains(&req["author"].as_i64().unwrap()) {
    109 					continue
    110 				}
    111                 if sink.send(Message::Text(msg)).await.is_err() {
    112                     break;
    113                 }
    114             }
    115         });
    116     })
    117 }
    118 
    119 #[tokio::main]
    120 async fn main() -> Result<(), std::io::Error> {
    121     if std::env::var_os("RUST_LOG").is_none() {
    122         std::env::set_var("RUST_LOG", "poem=debug");
    123     }
    124     tracing_subscriber::fmt::init();
    125 
    126     let app = Route::new().at(
    127         "/",
    128         get(ws.data(tokio::sync::broadcast::channel::<String>(32).0)),
    129     );
    130 
    131     Server::new(TcpListener::bind("127.0.0.1:3001")).run(app).await
    132 }
    133 
    134 // #[cfg(test)]
    135 // pub mod tests {
    136 //  use super::*;
    137 //  use websocket::{ClientBuilder, Message};
    138     
    139 //  #[tokio::test]
    140 //  async fn simple_messaging_flow() {              
    141 //      let sess = setup_db();      
    142 //         sess.execute(&stmt!(&format!(
    143 //             "CREATE KEYSPACE IF NOT EXISTS test \
    144 //              WITH replication = {{'class':'SimpleStrategy', 'replication_factor': 1}}"
    145 //         ))).wait().unwrap();
    146 //         sess.execute(&stmt!(&format!(
    147 //             "CREATE TABLE IF NOT EXISTS test.users \
    148 //              (id bigint PRIMARY KEY, name text, email text, hash text);"
    149 //         ))).wait().unwrap();
    150 //         sess.execute(&stmt!(&format!(
    151 //             "CREATE TABLE IF NOT EXISTS test.groups \
    152 //              (id bigint PRIMARY KEY, name text, members set<bigint>, is_dm boolean, \
    153 //              channels set<bigint>, admin set<bigint>, owner bigint);"
    154 //         ))).wait().unwrap();
    155 
    156 //      sess.execute(&stmt!(
    157 //          "INSERT INTO test.users (id, name, email, hash) VALUES (1234, 'steve', 'no@you.com', 'abc');"
    158 //      )).wait().unwrap();
    159 //      sess.execute(&stmt!(
    160 //          "INSERT INTO test.users (id, name, email, hash) VALUES (1235, 'erica', 'yes@you.com', 'abc3');"
    161 //      )).wait().unwrap();
    162 //      sess.execute(&stmt!(
    163 //          "INSERT INTO test.channels (id, group, name, members, private) VALUES (1111, 2222, 'main', {1234, 1235}, false);"
    164 //      )).wait().unwrap();
    165         
    166 //      // HACK HACK HACK
    167 //      std::process::Command::new("cargo")
    168 //             .args(["run", "-p", "chatterbox"])
    169 //             .spawn();
    170 //      std::thread::sleep(std::time::Duration::from_secs(10));
    171         
    172 //      let steve = std::thread::spawn(|| {                     
    173 //          let mut client = ClientBuilder::new("ws://127.0.0.1:3001")
    174 //              .unwrap()
    175 //              .connect_insecure()
    176 //              .unwrap();
    177             
    178 //          let message = Message::text("{\"hash\": \"abc\", \"id\": \"1234\"}");
    179 //          client.send_message(&message).unwrap();
    180 //          let message = Message::text("{\"content\": \"Hello\", \"channel\": \"1111\"}");
    181 //          client.send_message(&message).unwrap();
    182 //      });
    183 
    184 //      let erica = std::thread::spawn(|| {                     
    185 //          let mut client = ClientBuilder::new("ws://127.0.0.1:3001/")
    186 //              .unwrap()
    187 //              .connect_insecure()
    188 //              .unwrap();
    189             
    190 //          let message = Message::text("{\"hash\": \"abc3\", \"id\": \"1235\"}");
    191 //          client.send_message(&message).unwrap();
    192 //          let recv = client.recv_message().unwrap();
    193 //          if let websocket::OwnedMessage::Text(msg) = recv {
    194 //              let req: MessageObj = serde_json::from_str(&msg).unwrap();
    195 //              println!("{}", req.content);
    196 //              assert_eq!(req.author, 1234);
    197 //              assert_eq!(req.content, "Hello");
    198 //              assert_eq!(req.channel, 1111);
    199 //          } else {
    200 //              panic!("Got non-textual message!")
    201 //          }
    202 //      });
    203         
    204 //      steve.join().unwrap();
    205 //      erica.join().unwrap();
    206 //  }
    207 // }