316 lines
8.7 KiB
Rust
316 lines
8.7 KiB
Rust
use std::{net::SocketAddr, time::Duration};
|
|
|
|
use axum::{
|
|
Extension, Json, Router,
|
|
extract::{
|
|
ConnectInfo, Path, Query, WebSocketUpgrade,
|
|
ws::{Message as WsMessage, WebSocket},
|
|
},
|
|
http::{HeaderMap, StatusCode},
|
|
response::IntoResponse,
|
|
routing::{get, post},
|
|
};
|
|
use axum_extra::{TypedHeader, headers};
|
|
use sqlx::PgPool;
|
|
use tokio::select;
|
|
use uuid::Uuid;
|
|
|
|
use crate::{
|
|
auth::{verify_jwt, verify_jwt_string},
|
|
db::room_id_from_uuid,
|
|
errors::APIError,
|
|
routes::{rooms::is_member, ws::WsAuthQuery},
|
|
};
|
|
use crate::{
|
|
db::{user_id_from_uuid, username_from_uuid},
|
|
realtime::RealtimeMessages,
|
|
};
|
|
|
|
#[derive(sqlx::FromRow, serde::Serialize, Debug)]
|
|
pub struct MessageRow {
|
|
pub uuid: Uuid,
|
|
pub sender: String,
|
|
pub sender_uuid: Uuid,
|
|
pub room_uuid: Uuid,
|
|
pub message_type: String,
|
|
pub content: String,
|
|
pub sent_at: chrono::NaiveDateTime,
|
|
}
|
|
|
|
#[derive(serde::Serialize, Debug, Clone)]
|
|
pub struct Message {
|
|
pub uuid: Uuid,
|
|
pub room_uuid: Uuid,
|
|
pub sender: String,
|
|
pub sender_uuid: Uuid,
|
|
pub message_type: String,
|
|
pub content: String,
|
|
pub sent_at: String,
|
|
}
|
|
|
|
#[derive(serde::Deserialize)]
|
|
pub struct NewMessagePayload {
|
|
pub message_type: String,
|
|
pub content: String,
|
|
}
|
|
|
|
#[derive(serde::Deserialize)]
|
|
struct MessageFetchQuery {
|
|
limit: Option<i32>,
|
|
before: Option<Uuid>,
|
|
}
|
|
|
|
pub fn routes() -> Router {
|
|
Router::new()
|
|
.route("/messages/{room_uuid}", get(list_messages))
|
|
.route("/messages/{room_uuid}", post(create_message))
|
|
.route("/ws/messages", get(message_ws_handler))
|
|
}
|
|
|
|
/// Also resets `last_read_at`
|
|
async fn list_messages(
|
|
Path(room_uuid): Path<Uuid>,
|
|
Query(query): Query<MessageFetchQuery>,
|
|
headers: HeaderMap,
|
|
Extension(db): Extension<PgPool>,
|
|
) -> Result<Json<Vec<Message>>, APIError> {
|
|
let claims = verify_jwt(headers)?;
|
|
|
|
let user_id = user_id_from_uuid(&db, claims.sub).await?;
|
|
let room_id = room_id_from_uuid(&db, room_uuid).await?;
|
|
|
|
if !is_member(user_id, room_id, &db).await {
|
|
return Err(APIError::NotAMember);
|
|
}
|
|
|
|
let limit: i32 = query.limit.unwrap_or(30).abs().min(80);
|
|
|
|
let mut tx = db.begin().await?;
|
|
|
|
let messages = sqlx::query_as::<_, MessageRow>(
|
|
r#"
|
|
SELECT
|
|
m.uuid,
|
|
u.username AS sender,
|
|
u.uuid AS sender_uuid,
|
|
r.uuid AS room_uuid,
|
|
m.message_type,
|
|
m.content,
|
|
m.sent_at
|
|
FROM message_ m
|
|
JOIN user_ u ON u.id = m.sender
|
|
JOIN room_ r ON r.id = m.room
|
|
WHERE m.room = $1
|
|
AND ($2::uuid IS NULL OR m.id < (SELECT id FROM message_ WHERE uuid = $2))
|
|
ORDER BY m.id DESC
|
|
LIMIT $3
|
|
"#,
|
|
)
|
|
.bind(room_id)
|
|
.bind(query.before)
|
|
.bind(limit)
|
|
.fetch_all(&mut *tx)
|
|
.await?;
|
|
|
|
let mut messages: Vec<Message> = messages
|
|
.into_iter()
|
|
.map(|m| Message {
|
|
uuid: m.uuid,
|
|
room_uuid: m.room_uuid,
|
|
sender: m.sender,
|
|
sender_uuid: m.sender_uuid,
|
|
message_type: m.message_type,
|
|
content: m.content,
|
|
sent_at: m.sent_at.format("%Y-%m-%d %H:%M:%S").to_string(),
|
|
})
|
|
.collect();
|
|
|
|
messages.reverse();
|
|
|
|
sqlx::query(
|
|
r#"
|
|
UPDATE membership_
|
|
SET last_read_at = CURRENT_TIMESTAMP
|
|
WHERE user_id = $1
|
|
AND room = $2
|
|
"#,
|
|
)
|
|
.bind(user_id)
|
|
.bind(room_id)
|
|
.execute(&mut *tx)
|
|
.await?;
|
|
|
|
tx.commit().await?;
|
|
|
|
Ok(Json(messages))
|
|
}
|
|
|
|
async fn create_message(
|
|
Path(room_uuid): Path<Uuid>,
|
|
Extension(db): Extension<PgPool>,
|
|
Extension(realtime): Extension<RealtimeMessages>,
|
|
headers: HeaderMap,
|
|
Json(payload): Json<NewMessagePayload>,
|
|
) -> Result<(StatusCode, Json<Message>), APIError> {
|
|
let claims = verify_jwt(headers)?;
|
|
|
|
let user_id = user_id_from_uuid(&db, claims.sub).await?;
|
|
let room_id = room_id_from_uuid(&db, room_uuid).await?;
|
|
|
|
if !is_member(user_id, room_id, &db).await {
|
|
return Err(APIError::NotAMember);
|
|
}
|
|
|
|
let uuid = Uuid::now_v7();
|
|
|
|
let sent_at: chrono::NaiveDateTime = sqlx::query_scalar(
|
|
"INSERT INTO message_ (sender, room, message_type, content, uuid)
|
|
VALUES ($1, $2, $3, $4, $5) RETURNING sent_at",
|
|
)
|
|
.bind(user_id)
|
|
.bind(room_id)
|
|
.bind(&payload.message_type)
|
|
.bind(&payload.content)
|
|
.bind(uuid)
|
|
.fetch_one(&db)
|
|
.await?;
|
|
|
|
let sender_name = username_from_uuid(&db, claims.sub).await?;
|
|
|
|
let message = Message {
|
|
uuid,
|
|
room_uuid,
|
|
sender: sender_name,
|
|
sender_uuid: claims.sub,
|
|
message_type: payload.message_type,
|
|
content: payload.content,
|
|
sent_at: sent_at.format("%Y-%m-%d %H:%M:%S").to_string(),
|
|
};
|
|
|
|
let recipients: Vec<Uuid> = sqlx::query_scalar(
|
|
r#"
|
|
SELECT u.uuid
|
|
FROM membership_ m
|
|
JOIN user_ u ON u.id = m.user_id
|
|
WHERE m.room = $1
|
|
"#,
|
|
)
|
|
.bind(room_id)
|
|
.fetch_all(&db)
|
|
.await?;
|
|
|
|
let rt = realtime.clone();
|
|
let msg_clone = message.clone();
|
|
|
|
tokio::spawn(async move {
|
|
rt.broadcast(recipients, msg_clone);
|
|
});
|
|
|
|
Ok((StatusCode::CREATED, Json(message)))
|
|
}
|
|
|
|
async fn message_ws_handler(
|
|
ws: WebSocketUpgrade,
|
|
user_agent: Option<TypedHeader<headers::UserAgent>>,
|
|
Query(query): Query<WsAuthQuery>,
|
|
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
|
Extension(realtime): Extension<RealtimeMessages>,
|
|
Extension(db): Extension<sqlx::PgPool>,
|
|
) -> Result<impl IntoResponse, APIError> {
|
|
// tracing::info!("recieved ws handshake: {}", room_uuid);
|
|
|
|
let claims = verify_jwt_string(&query.token)?;
|
|
let user_uuid = claims.sub;
|
|
|
|
let result = sqlx::query(
|
|
r#"
|
|
delete from ws_token_
|
|
where token = $1
|
|
and expires_at > now()
|
|
"#,
|
|
)
|
|
.bind(query.token)
|
|
.execute(&db)
|
|
.await
|
|
// NOTE: Maybe wrong type of error
|
|
.map_err(|e| APIError::Internal(format!("Failed to get WS token from DB: {e}")))?;
|
|
|
|
if result.rows_affected() == 0 {
|
|
return Err(APIError::InvalidToken);
|
|
}
|
|
|
|
let receiver = realtime.get_sender(user_uuid).subscribe();
|
|
|
|
let user_agent = if let Some(TypedHeader(user_agent)) = user_agent {
|
|
user_agent.to_string()
|
|
} else {
|
|
String::from("Unknown browser")
|
|
};
|
|
tracing::debug!("`{user_agent}` {user_uuid} at {addr} connected.");
|
|
|
|
Ok(ws.on_upgrade(move |socket| handle_message_socket(socket, addr, receiver)))
|
|
}
|
|
|
|
async fn handle_message_socket(
|
|
mut socket: WebSocket,
|
|
who: SocketAddr,
|
|
mut receiver: tokio::sync::broadcast::Receiver<crate::routes::messages::Message>,
|
|
) {
|
|
let mut ping_interval = tokio::time::interval(Duration::from_secs(30));
|
|
|
|
loop {
|
|
select! {
|
|
// Receive broadcast messages and send to client (any room)
|
|
msg = receiver.recv() => {
|
|
if let Ok(msg) = msg {
|
|
if let Ok(json) = serde_json::to_string(&msg)
|
|
&& socket.send(WsMessage::Text(json.into())).await.is_err()
|
|
{
|
|
tracing::error!("Failed to send message to {who}, closing connection");
|
|
break;
|
|
}
|
|
} else {
|
|
break;
|
|
}
|
|
}
|
|
|
|
// Send Ping
|
|
_ = ping_interval.tick() => {
|
|
if socket.send(WsMessage::Ping(vec![].into())).await.is_err() {
|
|
tracing::error!("Failed to send ping to {who}, closing connection");
|
|
break;
|
|
}
|
|
|
|
// tracing::debug!("Ping sent to {who}");
|
|
}
|
|
|
|
// Get incoming messages from client
|
|
client_msg = socket.recv() => {
|
|
if let Some(Ok(msg)) = client_msg {
|
|
// match msg {
|
|
// // WsMessage::Pong(_) => {
|
|
// // tracing::debug!("Received Pong from {who}");
|
|
// // }
|
|
// // WsMessage::Ping(_) => {
|
|
// // tracing::info!("Received Ping from client");
|
|
// // }
|
|
// // WsMessage::Text(_) => {}
|
|
// WsMessage::Close(_) => {
|
|
// tracing::debug!("Client disconnected");
|
|
// break;
|
|
// }
|
|
// _ => {}
|
|
// }
|
|
if let WsMessage::Close(_) = msg {
|
|
tracing::debug!("Client disconnected");
|
|
break;
|
|
}
|
|
} else {
|
|
tracing::debug!("Client {who} abruptly disconnected");
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|