This commit is contained in:
parent
b613aeb9d9
commit
55de4a1fc0
11 changed files with 2105 additions and 141 deletions
37
src/entities/messages.rs
Normal file
37
src/entities/messages.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use sea_orm::entity::prelude::*;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Serialize, Deserialize)]
|
||||
#[sea_orm(table_name = "messages")]
|
||||
pub struct Model {
|
||||
#[sea_orm(primary_key, auto_increment = false)]
|
||||
pub id: Uuid,
|
||||
pub from_user: String,
|
||||
pub to_user: String,
|
||||
pub content: String,
|
||||
pub created_at: DateTime,
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
|
||||
pub enum Relation {
|
||||
#[sea_orm(
|
||||
belongs_to = "super::users::Entity",
|
||||
from = "Column::FromUser",
|
||||
to = "super::users::Column::Username"
|
||||
)]
|
||||
FromUser,
|
||||
#[sea_orm(
|
||||
belongs_to = "super::users::Entity",
|
||||
from = "Column::ToUser",
|
||||
to = "super::users::Column::Username"
|
||||
)]
|
||||
ToUser,
|
||||
}
|
||||
|
||||
impl Related<super::users::Entity> for Entity {
|
||||
fn to() -> RelationDef {
|
||||
Relation::FromUser.def() // Or ToUser, depending on what you want as the 'default' relationship
|
||||
}
|
||||
}
|
||||
|
||||
impl ActiveModelBehavior for ActiveModel {}
|
||||
7
src/entities/mod.rs
Normal file
7
src/entities/mod.rs
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
pub mod messages;
|
||||
pub mod users;
|
||||
|
||||
pub mod prelude {
|
||||
pub use super::messages::Entity as Messages;
|
||||
pub use super::users::Entity as Users;
|
||||
}
|
||||
26
src/entities/users.rs
Normal file
26
src/entities/users.rs
Normal file
|
|
@ -0,0 +1,26 @@
|
|||
use sea_orm::entity::prelude::*;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, DeriveEntityModel, Serialize, Deserialize)]
|
||||
#[sea_orm(table_name = "users")]
|
||||
pub struct Model {
|
||||
#[sea_orm(primary_key)]
|
||||
pub id: i32,
|
||||
#[sea_orm(unique)]
|
||||
pub username: String,
|
||||
pub created_at: DateTime,
|
||||
}
|
||||
|
||||
#[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)]
|
||||
pub enum Relation {
|
||||
#[sea_orm(has_many = "super::messages::Entity")]
|
||||
Messages,
|
||||
}
|
||||
|
||||
impl Related<super::messages::Entity> for Entity {
|
||||
fn to() -> RelationDef {
|
||||
Relation::Messages.def()
|
||||
}
|
||||
}
|
||||
|
||||
impl ActiveModelBehavior for ActiveModel {}
|
||||
191
src/main.rs
191
src/main.rs
|
|
@ -4,7 +4,7 @@ use axum::{
|
|||
Query, State,
|
||||
ws::{Message, WebSocket, WebSocketUpgrade},
|
||||
},
|
||||
http::StatusCode,
|
||||
http::{HeaderMap, StatusCode},
|
||||
response::{Html, IntoResponse},
|
||||
routing::get,
|
||||
};
|
||||
|
|
@ -15,7 +15,15 @@ use serde::{Deserialize, Serialize};
|
|||
use std::{collections::HashMap, sync::Arc};
|
||||
use tokio::sync::{RwLock, mpsc};
|
||||
|
||||
use migration::{Migrator, MigratorTrait};
|
||||
use sea_orm::{
|
||||
ActiveModelTrait, ActiveValue, ColumnTrait, Condition, Database, DatabaseConnection,
|
||||
EntityTrait, QueryFilter, QueryOrder,
|
||||
};
|
||||
|
||||
mod entities;
|
||||
mod user;
|
||||
use entities::{messages, prelude::*, users};
|
||||
use user::User;
|
||||
|
||||
// WebSocket message types for client-server communication
|
||||
|
|
@ -23,13 +31,18 @@ use user::User;
|
|||
#[serde(tag = "type")]
|
||||
enum WsMessage {
|
||||
#[serde(rename = "user_list")]
|
||||
UserList { users: Vec<String> },
|
||||
UserList {
|
||||
users: Vec<String>,
|
||||
online: Vec<String>,
|
||||
},
|
||||
#[serde(rename = "private_message")]
|
||||
PrivateMessage {
|
||||
from: String,
|
||||
to: String,
|
||||
content: String,
|
||||
},
|
||||
#[serde(rename = "identity")]
|
||||
Identity { username: String },
|
||||
#[serde(rename = "user_joined")]
|
||||
UserJoined { username: String },
|
||||
#[serde(rename = "system")]
|
||||
|
|
@ -90,16 +103,27 @@ impl ChatState {
|
|||
struct AppState {
|
||||
chat: ChatState,
|
||||
jwks: RwLock<HashMap<String, (String, String)>>,
|
||||
db: DatabaseConnection,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
dotenv().ok();
|
||||
|
||||
let database_url = dotenvy::var("DATABASE_URL").expect("DATABASE_URL must be set");
|
||||
let db = Database::connect(database_url)
|
||||
.await
|
||||
.expect("Failed to connect to database");
|
||||
|
||||
// Run migrations
|
||||
Migrator::up(&db, None)
|
||||
.await
|
||||
.expect("Failed to run migrations");
|
||||
|
||||
let chat = ChatState::new();
|
||||
let jwks = RwLock::new(HashMap::new());
|
||||
|
||||
let app_state = Arc::new(AppState { chat, jwks });
|
||||
let app_state = Arc::new(AppState { chat, jwks, db });
|
||||
|
||||
// Initial JWKS fetch
|
||||
if let Err(e) = fetch_jwks(&app_state).await {
|
||||
|
|
@ -109,6 +133,7 @@ async fn main() {
|
|||
// Build application with routes
|
||||
let app = Router::new()
|
||||
.route("/", get(index))
|
||||
.route("/api/history", get(get_history))
|
||||
.route("/ws", get(websocket_handler))
|
||||
.with_state(app_state);
|
||||
|
||||
|
|
@ -166,19 +191,15 @@ struct Claims {
|
|||
// Add other fields as needed
|
||||
}
|
||||
|
||||
async fn websocket_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
Query(params): Query<WsParams>,
|
||||
State(state): State<Arc<AppState>>,
|
||||
) -> impl IntoResponse {
|
||||
let header = match decode_header(¶ms.token) {
|
||||
async fn verify_token(token: &str, state: &AppState) -> Result<Claims, StatusCode> {
|
||||
let header = match decode_header(token) {
|
||||
Ok(h) => h,
|
||||
Err(_) => return StatusCode::BAD_REQUEST.into_response(),
|
||||
Err(_) => return Err(StatusCode::BAD_REQUEST),
|
||||
};
|
||||
|
||||
let kid = match header.kid {
|
||||
Some(k) => k,
|
||||
None => return StatusCode::UNAUTHORIZED.into_response(),
|
||||
None => return Err(StatusCode::UNAUTHORIZED),
|
||||
};
|
||||
|
||||
let components = {
|
||||
|
|
@ -188,44 +209,123 @@ async fn websocket_handler(
|
|||
} else {
|
||||
// Key not found, might need to refresh JWKS
|
||||
drop(jwks);
|
||||
if let Err(e) = fetch_jwks(&state).await {
|
||||
if let Err(e) = fetch_jwks(state).await {
|
||||
eprintln!("Failed to refresh JWKS: {}", e);
|
||||
return StatusCode::UNAUTHORIZED.into_response();
|
||||
return Err(StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
let jwks = state.jwks.read().await;
|
||||
match jwks.get(&kid).cloned() {
|
||||
Some(c) => c,
|
||||
None => return StatusCode::UNAUTHORIZED.into_response(),
|
||||
None => return Err(StatusCode::UNAUTHORIZED),
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let key = match DecodingKey::from_rsa_components(&components.0, &components.1) {
|
||||
Ok(k) => k,
|
||||
Err(_) => return StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
||||
Err(_) => return Err(StatusCode::INTERNAL_SERVER_ERROR),
|
||||
};
|
||||
|
||||
let mut validation = Validation::new(Algorithm::RS256);
|
||||
validation.set_audience(&["xLL8a23JBnBvEarjQUY6Jgd4zLnnqIi2y3oB6laM"]);
|
||||
validation.set_issuer(&["https://idm.flegr.me/application/o/chat/"]);
|
||||
|
||||
let token_data = match decode::<Claims>(¶ms.token, &key, &validation) {
|
||||
Ok(c) => c,
|
||||
match decode::<Claims>(token, &key, &validation) {
|
||||
Ok(c) => Ok(c.claims),
|
||||
Err(e) => {
|
||||
eprintln!("Token validation failed: {}", e);
|
||||
return StatusCode::UNAUTHORIZED.into_response();
|
||||
Err(StatusCode::UNAUTHORIZED)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn websocket_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
Query(params): Query<WsParams>,
|
||||
State(state): State<Arc<AppState>>,
|
||||
) -> impl IntoResponse {
|
||||
let claims = match verify_token(¶ms.token, &state).await {
|
||||
Ok(c) => c,
|
||||
Err(status) => return status.into_response(),
|
||||
};
|
||||
|
||||
let user = User {
|
||||
login: token_data.claims.preferred_username,
|
||||
login: claims.preferred_username,
|
||||
avatar_url: String::new(),
|
||||
};
|
||||
|
||||
// Ensure user exists in DB
|
||||
let user_active_model = users::ActiveModel {
|
||||
username: ActiveValue::Set(user.login.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
let _ = Users::insert(user_active_model)
|
||||
.on_conflict(
|
||||
sea_orm::sea_query::OnConflict::column(users::Column::Username)
|
||||
.do_nothing()
|
||||
.to_owned(),
|
||||
)
|
||||
.exec(&state.db)
|
||||
.await;
|
||||
|
||||
let state_clone = state.clone();
|
||||
ws.on_upgrade(move |socket| handle_websocket(socket, state_clone, user))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct HistoryParams {
|
||||
to: String,
|
||||
}
|
||||
|
||||
async fn get_history(
|
||||
headers: HeaderMap,
|
||||
Query(params): Query<HistoryParams>,
|
||||
State(state): State<Arc<AppState>>,
|
||||
) -> impl IntoResponse {
|
||||
let auth_header = match headers.get(axum::http::header::AUTHORIZATION) {
|
||||
Some(h) => h.to_str().unwrap_or(""),
|
||||
None => return StatusCode::UNAUTHORIZED.into_response(),
|
||||
};
|
||||
|
||||
if !auth_header.starts_with("Bearer ") {
|
||||
return StatusCode::UNAUTHORIZED.into_response();
|
||||
}
|
||||
|
||||
let token = &auth_header[7..];
|
||||
let claims = match verify_token(token, &state).await {
|
||||
Ok(c) => c,
|
||||
Err(status) => return status.into_response(),
|
||||
};
|
||||
|
||||
let current_user = claims.preferred_username;
|
||||
|
||||
let messages = Messages::find()
|
||||
.filter(
|
||||
Condition::any()
|
||||
.add(
|
||||
Condition::all()
|
||||
.add(messages::Column::FromUser.eq(current_user.clone()))
|
||||
.add(messages::Column::ToUser.eq(params.to.clone())),
|
||||
)
|
||||
.add(
|
||||
Condition::all()
|
||||
.add(messages::Column::FromUser.eq(params.to))
|
||||
.add(messages::Column::ToUser.eq(current_user)),
|
||||
),
|
||||
)
|
||||
.order_by_asc(messages::Column::CreatedAt)
|
||||
.all(&state.db)
|
||||
.await;
|
||||
|
||||
match messages {
|
||||
Ok(msgs) => axum::Json(msgs).into_response(),
|
||||
Err(e) => {
|
||||
eprintln!("Failed to fetch history: {}", e);
|
||||
StatusCode::INTERNAL_SERVER_ERROR.into_response()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WebSocket connection handler
|
||||
async fn handle_websocket(stream: WebSocket, state: Arc<AppState>, user: User) {
|
||||
let (mut ws_sender, mut ws_receiver) = stream.split();
|
||||
|
|
@ -238,18 +338,41 @@ async fn handle_websocket(stream: WebSocket, state: Arc<AppState>, user: User) {
|
|||
state.chat.add_user(username.clone(), tx).await;
|
||||
println!("User connected: {}", username);
|
||||
|
||||
// Send current user list to the newly connected user
|
||||
let users = state.chat.get_online_users().await;
|
||||
println!("Sending user list to {}: {:?}", username, users);
|
||||
let user_list_msg = serde_json::to_string(&WsMessage::UserList { users }).unwrap();
|
||||
let _ = ws_sender.send(Message::Text(user_list_msg.into())).await;
|
||||
|
||||
// Send welcome message
|
||||
let welcome = serde_json::to_string(&WsMessage::System {
|
||||
content: format!("Welcome, {}!", username),
|
||||
// Send identity
|
||||
let identity_msg = serde_json::to_string(&WsMessage::Identity {
|
||||
username: username.clone(),
|
||||
})
|
||||
.unwrap();
|
||||
let _ = ws_sender.send(Message::Text(welcome.into())).await;
|
||||
let _ = ws_sender.send(Message::Text(identity_msg.into())).await;
|
||||
|
||||
// Send global user list and online status to the newly connected user
|
||||
let (db_users, online_users) = match Users::find().all(&state.db).await {
|
||||
Ok(all_users) => {
|
||||
let db_list = all_users
|
||||
.into_iter()
|
||||
.map(|u| u.username)
|
||||
.filter(|u| u != &username)
|
||||
.collect::<Vec<String>>();
|
||||
let online_list = state.chat.get_online_users().await;
|
||||
(db_list, online_list)
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("Failed to fetch users from DB: {}", e);
|
||||
(Vec::new(), Vec::new())
|
||||
}
|
||||
};
|
||||
println!(
|
||||
"Sending global user list to {}: DB Count={}, Online Count={}",
|
||||
username,
|
||||
db_users.len(),
|
||||
online_users.len()
|
||||
);
|
||||
let user_list_msg = serde_json::to_string(&WsMessage::UserList {
|
||||
users: db_users,
|
||||
online: online_users,
|
||||
})
|
||||
.unwrap();
|
||||
let _ = ws_sender.send(Message::Text(user_list_msg.into())).await;
|
||||
|
||||
// Task to forward messages from channel to WebSocket
|
||||
let send_task = tokio::spawn(async move {
|
||||
|
|
@ -274,6 +397,16 @@ async fn handle_websocket(stream: WebSocket, state: Arc<AppState>, user: User) {
|
|||
while let Some(Ok(Message::Text(text))) = ws_receiver.next().await {
|
||||
// Parse incoming message
|
||||
if let Ok(msg) = serde_json::from_str::<ClientMessage>(&text) {
|
||||
// Persist to DB
|
||||
let new_message = messages::ActiveModel {
|
||||
id: ActiveValue::Set(uuid::Uuid::new_v4()),
|
||||
from_user: ActiveValue::Set(username_clone.clone()),
|
||||
to_user: ActiveValue::Set(msg.to.clone()),
|
||||
content: ActiveValue::Set(msg.content.clone()),
|
||||
..Default::default()
|
||||
};
|
||||
let _ = new_message.insert(&state_clone.db).await;
|
||||
|
||||
// Create private message
|
||||
let private_msg = serde_json::to_string(&WsMessage::PrivateMessage {
|
||||
from: username_clone.clone(),
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue