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 // }