This commit is contained in:
2026-08-09 10:30:54 +02:00
parent 93800e8460
commit 20beea24d5
11 changed files with 325 additions and 61 deletions
+53 -4
View File
@@ -1,17 +1,19 @@
use crate::core::AppState;
use crate::http::context::CurrentUser;
use crate::models::user::Model as User;
use crate::permissions::ChannelPermission;
use crate::routes::gateway::GatewayClient;
use axum::{
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
State,
ws::{Message, WebSocket, WebSocketUpgrade},
},
response::IntoResponse,
};
use futures_util::{sink::SinkExt, stream::StreamExt};
use serde::Deserialize;
use tokio::sync::mpsc;
use uuid::Uuid;
#[derive(Deserialize)]
pub struct WsQuery {
@@ -41,9 +43,17 @@ async fn handle_socket(socket: WebSocket, state: AppState, user: User) {
let (mut sender, mut receiver) = socket.split();
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let channel_ids = match accessible_channel_ids(&state, &user).await {
Ok(channel_ids) => channel_ids,
Err(error) => {
tracing::error!(user_id = %user.id, ?error, "Unable to resolve gateway channel access");
return;
}
};
let mut client = GatewayClient::new(user, tx, state.event_bus.clone());
client.on_connect().await;
state.gateway.add_client(client.clone());
state.gateway.add_client(client.clone(), channel_ids);
// Task pour envoyer les messages du canal mpsc vers le WebSocket
let mut send_task = tokio::spawn(async move {
@@ -56,7 +66,6 @@ async fn handle_socket(socket: WebSocket, state: AppState, user: User) {
// Task pour recevoir les messages du WebSocket
let client_clone = client.clone();
let state_clone = state.clone();
let mut recv_task = tokio::spawn(async move {
while let Some(Ok(message)) = receiver.next().await {
client_clone.on_message(message).await;
@@ -69,7 +78,47 @@ async fn handle_socket(socket: WebSocket, state: AppState, user: User) {
_ = (&mut recv_task) => send_task.abort(),
};
state.gateway.remove_client(client.clone());
state.gateway.remove_client(&client);
// // Déconnexion (Disconnect)
client.on_disconnect().await;
}
async fn accessible_channel_ids(state: &AppState, user: &User) -> Result<Vec<Uuid>, anyhow::Error> {
if user.is_superuser {
return Ok(state
.repositories
.channel
.get_all()
.await?
.into_iter()
.map(|channel| channel.id)
.collect());
}
let mut channel_ids = Vec::new();
for server in state.repositories.server.get_all().await? {
if state
.repositories
.server
.get_user(server.id, user.id)
.await?
.is_none()
{
continue;
}
let tree = state
.repositories
.server_tree
.get_for_user(server.id, user.id)
.await?;
channel_ids.extend(tree.channels.into_iter().filter_map(|channel| {
channel
.permissions
.filter(|permissions| permissions.contains(ChannelPermission::READ_CHANNEL))
.map(|_| channel.channel.id)
}));
}
Ok(channel_ids)
}
+117 -42
View File
@@ -1,17 +1,16 @@
use crate::models::category;
use crate::models::channel;
use crate::models::message;
use crate::models::server;
use crate::domain::events::message::{
MessageCreatedEvent, MessageDeletedEvent, MessageUpdatedEvent,
};
use crate::models::user::Model as User;
use crate::routes::category::mapper::category_model_to_category_response;
use crate::routes::channel::mapper::channel_model_to_channel_response;
use crate::routes::message::mapper::message_model_to_message_response;
use crate::routes::message::mapper::message_model_to_message_response_with_server_id;
use crate::routes::server::mapper::server_model_to_server_response;
use axum::extract::ws::Message;
use event_bus::EventBus;
use events::GatewayEvent;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
@@ -23,8 +22,15 @@ pub mod routes;
#[derive(Debug, Default)]
pub struct GatewayManager {
// {UserID: {connection_id: GatewayClient}}
pub clients: RwLock<HashMap<Uuid, HashMap<Uuid, GatewayClient>>>,
// Chaque connexion est inscrite dans les groupes des canaux accessibles.
pub clients: RwLock<HashMap<ConnectionKey, GatewayClient>>,
channel_subscribers: RwLock<HashMap<Uuid, HashSet<ConnectionKey>>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct ConnectionKey {
pub user_id: Uuid,
pub connection_id: Uuid,
}
#[derive(Debug, Clone)]
@@ -37,19 +43,98 @@ pub struct GatewayClient {
}
impl GatewayManager {
fn add_client(&self, gateway_client: GatewayClient) {
let mut clients = self.clients.write();
let user_id = gateway_client.user.id;
clients
.entry(user_id)
.or_insert_with(HashMap::new)
.insert(gateway_client.connection_id, gateway_client);
/// Démarre les routeurs centraux des événements de messages.
pub fn start(self: &Arc<Self>, event_bus: Arc<EventBus>) {
let manager = Arc::clone(self);
event_bus.on_async::<MessageCreatedEvent, _, _>("message_created", move |event| {
let manager = Arc::clone(&manager);
async move {
manager.broadcast_message(
event.channel_id,
"add",
message_model_to_message_response_with_server_id(
event.message,
event.server_id,
),
);
}
});
let manager = Arc::clone(self);
event_bus.on_async::<MessageUpdatedEvent, _, _>("message_updated", move |event| {
let manager = Arc::clone(&manager);
async move {
manager.broadcast_message(
event.channel_id,
"update",
message_model_to_message_response_with_server_id(
event.message,
event.server_id,
),
);
}
});
let manager = Arc::clone(self);
event_bus.on_async::<MessageDeletedEvent, _, _>("message_deleted", move |event| {
let manager = Arc::clone(&manager);
async move {
manager.broadcast_message(event.channel_id, "remove", event.message.id);
}
});
}
fn remove_client(&self, gateway_client: GatewayClient) {
let mut clients = self.clients.write();
if let Some(client_list) = clients.get_mut(&gateway_client.user.id) {
client_list.remove(&gateway_client.connection_id);
pub(crate) fn add_client(
&self,
gateway_client: GatewayClient,
channel_ids: impl IntoIterator<Item = Uuid>,
) {
let key = gateway_client.key();
self.clients.write().insert(key, gateway_client);
let mut subscribers = self.channel_subscribers.write();
for channel_id in channel_ids {
subscribers.entry(channel_id).or_default().insert(key);
}
}
pub(crate) fn remove_client(&self, gateway_client: &GatewayClient) {
let key = gateway_client.key();
self.clients.write().remove(&key);
let mut subscribers = self.channel_subscribers.write();
for channel_subscribers in subscribers.values_mut() {
channel_subscribers.remove(&key);
}
subscribers.retain(|_, values| !values.is_empty());
}
fn broadcast_message<T: serde::Serialize>(
&self,
channel_id: Uuid,
action: &'static str,
content: T,
) {
let event = GatewayEvent {
namespace: "Message",
action,
content,
};
let Ok(json) = serde_json::to_string(&event) else {
return;
};
let keys = self
.channel_subscribers
.read()
.get(&channel_id)
.cloned()
.unwrap_or_default();
let clients = self.clients.read();
for key in keys {
if let Some(client) = clients.get(&key) {
let _ = client.sender.send(Message::Text(json.clone().into()));
}
}
}
}
@@ -60,16 +145,22 @@ impl GatewayClient {
sender: mpsc::UnboundedSender<Message>,
event_bus: Arc<EventBus>,
) -> Self {
let connection_id = Uuid::new_v4();
Self {
user,
connection_id,
connection_id: Uuid::new_v4(),
sender,
event_bus,
_event_handles: Vec::new(),
}
}
pub fn key(&self) -> ConnectionKey {
ConnectionKey {
user_id: self.user.id,
connection_id: self.connection_id,
}
}
fn subscribe_event<T, F, R>(
&self,
event_name: &'static str,
@@ -105,22 +196,8 @@ impl GatewayClient {
pub fn subscribe_to_events(&mut self) {
let mut handles = Vec::new();
// Message
handles.push(self.subscribe_event(
"message_created",
"Message",
"add",
message_model_to_message_response,
));
handles.push(self.subscribe_event(
"message_updated",
"Message",
"update",
message_model_to_message_response,
));
handles.push(self.subscribe_event("message_deleted", "Message", "remove", |id: Uuid| id));
// Les messages sont routés par GatewayManager selon le channel_id.
// Channel
handles.push(self.subscribe_event(
"channel_created",
"Channel",
@@ -135,7 +212,6 @@ impl GatewayClient {
));
handles.push(self.subscribe_event("channel_deleted", "Channel", "remove", |id: Uuid| id));
// Category
handles.push(self.subscribe_event(
"category_created",
"Category",
@@ -150,7 +226,6 @@ impl GatewayClient {
));
handles.push(self.subscribe_event("category_deleted", "Category", "remove", |id: Uuid| id));
// Server
handles.push(self.subscribe_event(
"server_created",
"Server",
@@ -175,20 +250,20 @@ impl GatewayClient {
}
async fn on_connect(&mut self) {
tracing::info!("Client connected: {:?}", self.user);
tracing::info!(user_id = %self.user.id, "Client connected");
self.subscribe_to_events();
}
async fn on_disconnect(&mut self) {
tracing::info!("Client disconnected: {:?}", self.user);
tracing::info!(user_id = %self.user.id, "Client disconnected");
self.unsubscribe_all();
}
async fn on_message(&self, message: Message) {
match message {
Message::Binary(content) => {}
Message::Binary(_) => {}
Message::Text(content) => {
tracing::info!("Received text message: {}", content);
tracing::info!(user_id = %self.user.id, "Received text message: {}", content);
}
Message::Ping(_) => {}
Message::Pong(_) => {}