Working auth and fixed a lot of bugs
This commit is contained in:
+8
-8
@@ -6,15 +6,15 @@ use crate::{Server, utils::client::Client};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AuthApiRes {
|
||||
uuid: u32,
|
||||
user_id: u32,
|
||||
}
|
||||
|
||||
pub fn auth(_server: &Arc<Server>, client: &mut Client, token: &str) -> crate::Result<u32> {
|
||||
// let mut res = ureq::get(format!("https://api.voxa.org/server-auth?token={token}")).call()?;
|
||||
|
||||
// let api_res: AuthApiRes = serde_json::from_str(&res.body_mut().read_to_string()?)?;
|
||||
|
||||
client.set_uuid(232);
|
||||
|
||||
Ok(232)
|
||||
let mut res = ureq::get(format!(
|
||||
"http://localhost:3000/api/auth?intents=server&token={token}"
|
||||
))
|
||||
.call()?;
|
||||
let api_res: AuthApiRes = serde_json::from_str(&res.body_mut().read_to_string()?)?;
|
||||
client.set_uuid(api_res.user_id);
|
||||
Ok(api_res.user_id)
|
||||
}
|
||||
|
||||
+28
-10
@@ -89,7 +89,8 @@ impl Server {
|
||||
Ok(stream) => {
|
||||
std::thread::spawn({
|
||||
let srv = self.clone();
|
||||
move || match srv.handle_client(stream) {
|
||||
let client = srv.init_client(stream)?;
|
||||
move || match srv.wrap_err(&client, srv.handle_client(&client)) {
|
||||
Ok(_) => {}
|
||||
Err(e) => Self::LOGGER.error(format!("Client handler failed: {e}")),
|
||||
}
|
||||
@@ -108,7 +109,7 @@ impl Server {
|
||||
self.plugins.lock().unwrap().push(plugin);
|
||||
}
|
||||
|
||||
fn handle_client(self: &Arc<Self>, stream: TcpStream) -> anyhow::Result<()> {
|
||||
fn init_client(self: &Arc<Self>, stream: TcpStream) -> anyhow::Result<Client> {
|
||||
Self::LOGGER.info(format!("New connection: {}", stream.peer_addr()?));
|
||||
// Initialize client
|
||||
let mut client = Client::new(stream)?;
|
||||
@@ -125,20 +126,29 @@ impl Server {
|
||||
|
||||
match self.wrap_err(&client, client.read_t::<types::handshake::ClientDetails>())? {
|
||||
Some(types::WsMessage::Message(types::handshake::ClientDetails {
|
||||
auth_token, ..
|
||||
auth_token,
|
||||
last_message,
|
||||
..
|
||||
})) => {
|
||||
let auth_res = auth::auth(self, &mut client, &auth_token);
|
||||
let uuid = self.wrap_err(&client, auth_res)?;
|
||||
self.wrap_err(
|
||||
&client,
|
||||
client.send(types::ServerMessage::Authenticated { uuid }),
|
||||
client.send(types::ServerMessage::Authenticated {
|
||||
uuid,
|
||||
messages: if let Some(i) = last_message {
|
||||
self.wrap_err(&client, self.db.get_messages_after_id(i))?
|
||||
} else {
|
||||
self.wrap_err(&client, self.db.get_messages_after_id(0))?
|
||||
},
|
||||
}),
|
||||
)?;
|
||||
}
|
||||
Some(_) => {
|
||||
Some(v) => {
|
||||
self.wrap_err(
|
||||
&client,
|
||||
client.send(types::ResponseError::InvalidHandshake(format!(
|
||||
"Invalid handshake"
|
||||
"Invalid handshake: {v:?}"
|
||||
))),
|
||||
)?;
|
||||
}
|
||||
@@ -148,29 +158,37 @@ impl Server {
|
||||
// Insert to the set of all connected clients
|
||||
self.clients.lock().unwrap().insert(client.clone());
|
||||
|
||||
Ok(client)
|
||||
}
|
||||
|
||||
fn handle_client(self: &Arc<Self>, client: &Client) -> anyhow::Result<()> {
|
||||
// The main req/res loop
|
||||
'outer: loop {
|
||||
let req = client.read()?;
|
||||
if let Some(r) = &req {
|
||||
for p in self.plugins.lock().unwrap().iter_mut() {
|
||||
if p.on_request(r, &client, self) {
|
||||
if p.on_request(r, client, self) {
|
||||
continue 'outer;
|
||||
}
|
||||
}
|
||||
|
||||
self.call_request(r, &client)?;
|
||||
self.wrap_err(&client, self.call_request(r, &client))?;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// When there is a error it removes the client
|
||||
pub fn wrap_err<T, E>(
|
||||
pub fn wrap_err<T, E: std::fmt::Display>(
|
||||
self: &Arc<Self>,
|
||||
client: &Client,
|
||||
res: std::result::Result<T, E>,
|
||||
) -> std::result::Result<T, E> {
|
||||
if res.is_err() {
|
||||
if let Err(e) = &res {
|
||||
self.clients.lock().unwrap().remove(&client);
|
||||
if client
|
||||
.send(types::ResponseError::InternalError(e.to_string()))
|
||||
.is_err()
|
||||
{}
|
||||
}
|
||||
|
||||
res
|
||||
|
||||
+19
-15
@@ -13,27 +13,31 @@ pub fn send(
|
||||
LOGGER.info(format!("SendMessage to {channel_id}: {contents}"));
|
||||
|
||||
if contents.is_empty() {
|
||||
server.wrap_err(
|
||||
&client,
|
||||
client.send(types::ResponseError::InvalidRequest(format!(
|
||||
"Invalid message: empty message"
|
||||
))),
|
||||
)?;
|
||||
client.send(types::ResponseError::InvalidRequest(format!(
|
||||
"Invalid message: empty message"
|
||||
)))?;
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let msg = server.wrap_err(
|
||||
&client,
|
||||
server.db.insert_message(
|
||||
&channel_id,
|
||||
client.get_uuid()?,
|
||||
&contents,
|
||||
chrono::Utc::now().timestamp(),
|
||||
),
|
||||
let msg = server.db.insert_message(
|
||||
&channel_id,
|
||||
client.get_uuid()?,
|
||||
&contents,
|
||||
chrono::Utc::now().timestamp(),
|
||||
)?;
|
||||
|
||||
let server = server.clone();
|
||||
|
||||
for c in server.clients.lock().unwrap().iter() {
|
||||
server.wrap_err(&c, c.send(types::ServerMessage::MessageCreate(msg.clone())))?;
|
||||
let c = c.clone();
|
||||
let server = server.clone();
|
||||
let msg = msg.clone();
|
||||
std::thread::spawn(move || {
|
||||
server
|
||||
.wrap_err(&c, c.send(types::ServerMessage::MessageCreate(msg)))
|
||||
.expect("Failed to broadcast");
|
||||
});
|
||||
}
|
||||
|
||||
Ok(())
|
||||
|
||||
@@ -27,6 +27,7 @@ pub enum ServerMessage {
|
||||
/// Successful authentication
|
||||
Authenticated {
|
||||
uuid: u32,
|
||||
messages: Vec<data::Message>,
|
||||
},
|
||||
|
||||
TempMessage {
|
||||
@@ -118,5 +119,6 @@ pub mod handshake {
|
||||
pub struct ClientDetails {
|
||||
pub version: String,
|
||||
pub auth_token: String,
|
||||
pub last_message: Option<usize>,
|
||||
}
|
||||
}
|
||||
|
||||
+5
-2
@@ -341,7 +341,7 @@ impl Clone for Client {
|
||||
|
||||
impl PartialEq for Client {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.0.peer_addr().unwrap() == other.0.peer_addr().unwrap()
|
||||
self.0.peer_addr().unwrap_or(self.0.local_addr().unwrap()) == other.0.peer_addr().unwrap_or(other.0.local_addr().unwrap())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -349,6 +349,9 @@ impl Eq for Client {}
|
||||
|
||||
impl Hash for Client {
|
||||
fn hash<H: Hasher>(&self, state: &mut H) {
|
||||
self.0.peer_addr().unwrap().hash(state);
|
||||
self.0
|
||||
.peer_addr()
|
||||
.unwrap_or(self.0.local_addr().unwrap())
|
||||
.hash(state);
|
||||
}
|
||||
}
|
||||
|
||||
+28
-1
@@ -12,7 +12,7 @@ impl Database {
|
||||
"CREATE TABLE IF NOT EXISTS chat (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
channel_id TEXT NOT NULL,
|
||||
user_id TEXT NOT NULL,
|
||||
user_id INTEGER NOT NULL,
|
||||
contents TEXT NOT NULL,
|
||||
timestamp INTEGER NOT NULL
|
||||
)",
|
||||
@@ -102,6 +102,33 @@ impl Database {
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// Get all messages with an ID greater than the given one
|
||||
pub fn get_messages_after_id(&self, message_id: usize) -> Result<Vec<Message>> {
|
||||
let mut stmt = self.0.prepare(
|
||||
"SELECT id, channel_id, user_id, contents, timestamp
|
||||
FROM chat
|
||||
WHERE id > ?1
|
||||
ORDER BY id ASC",
|
||||
)?;
|
||||
|
||||
let rows = stmt.query_map(params![message_id], |row| {
|
||||
Ok(Message {
|
||||
id: row.get::<_, i64>(0)?,
|
||||
channel_id: row.get::<_, String>(1)?,
|
||||
from: row.get::<_, u32>(2)?,
|
||||
contents: row.get::<_, String>(3)?,
|
||||
timestamp: row.get::<_, i64>(4)?,
|
||||
})
|
||||
})?;
|
||||
|
||||
let mut messages = Vec::new();
|
||||
for row in rows {
|
||||
messages.push(row?);
|
||||
}
|
||||
|
||||
Ok(messages)
|
||||
}
|
||||
}
|
||||
|
||||
unsafe impl Send for Database {}
|
||||
|
||||
Reference in New Issue
Block a user