use axum::{
extract::{State, WebSocketUpgrade},
response::Response,
routing::get,
Router,
};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use uuid::Uuid;
pub mod client;
pub mod handler;
pub use handler::WebSocketHandler;
#[derive(Clone)]
pub struct RdtState {
store: Arc<crate::document::DocumentStore>,
clients: Arc<ClientManager>,
}
impl RdtState {
pub fn new(store: Arc<crate::document::DocumentStore>) -> Self {
Self {
store,
clients: Arc::new(ClientManager::new()),
}
}
pub fn store(&self) -> &Arc<crate::document::DocumentStore> {
&self.store
}
pub fn clients(&self) -> &Arc<ClientManager> {
&self.clients
}
}
pub struct ClientManager {
clients: RwLock<HashMap<String, ClientInfo>>,
subscriptions: RwLock<HashMap<(String, String), Vec<String>>>,
}
pub struct ClientInfo {
pub id: String,
pub sender: tokio::sync::mpsc::UnboundedSender<crate::protocol::ServerMessage>,
}
impl ClientManager {
pub fn new() -> Self {
Self {
clients: RwLock::new(HashMap::new()),
subscriptions: RwLock::new(HashMap::new()),
}
}
pub async fn register_client(
&self,
sender: tokio::sync::mpsc::UnboundedSender<crate::protocol::ServerMessage>,
) -> String {
let client_id = Uuid::new_v4().to_string();
let client_info = ClientInfo {
id: client_id.clone(),
sender,
};
self.clients
.write()
.await
.insert(client_id.clone(), client_info);
tracing::info!("Registered client: {}", client_id);
client_id
}
pub async fn unregister_client(&self, client_id: &str) {
self.clients.write().await.remove(client_id);
let mut subscriptions = self.subscriptions.write().await;
subscriptions.retain(|_, client_ids| {
client_ids.retain(|id| id != client_id);
!client_ids.is_empty()
});
tracing::info!("Unregistered client: {}", client_id);
}
pub async fn subscribe_client(&self, client_id: &str, document_id: &str, map_key: &str) {
let key = (document_id.to_string(), map_key.to_string());
let mut subscriptions = self.subscriptions.write().await;
subscriptions
.entry(key)
.or_insert_with(Vec::new)
.push(client_id.to_string());
tracing::debug!(
"Client {} subscribed to document '{}', map '{}'",
client_id,
document_id,
map_key
);
}
pub async fn unsubscribe_client(&self, client_id: &str, document_id: &str, map_key: &str) {
let key = (document_id.to_string(), map_key.to_string());
let mut subscriptions = self.subscriptions.write().await;
if let Some(client_ids) = subscriptions.get_mut(&key) {
client_ids.retain(|id| id != client_id);
if client_ids.is_empty() {
subscriptions.remove(&key);
}
}
tracing::debug!(
"Client {} unsubscribed from document '{}', map '{}'",
client_id,
document_id,
map_key
);
}
pub async fn broadcast_to_subscribers(
&self,
document_id: &str,
map_key: &str,
message: crate::protocol::ServerMessage,
) {
let key = (document_id.to_string(), map_key.to_string());
let subscriptions = self.subscriptions.read().await;
if let Some(client_ids) = subscriptions.get(&key) {
let clients = self.clients.read().await;
let message_arc = std::sync::Arc::new(message);
for client_id in client_ids {
if let Some(client_info) = clients.get(client_id) {
if client_info.sender.send((*message_arc).clone()).is_err() {
tracing::warn!("Failed to send message to client {}", client_id);
}
}
}
}
}
pub async fn send_to_client(&self, client_id: &str, message: crate::protocol::ServerMessage) {
let clients = self.clients.read().await;
if let Some(client_info) = clients.get(client_id) {
if client_info.sender.send(message).is_err() {
tracing::warn!("Failed to send message to client {}", client_id);
}
}
}
}
impl Default for ClientManager {
fn default() -> Self {
Self::new()
}
}
pub fn router_with_rdt(store: Arc<crate::document::DocumentStore>) -> Router<RdtState> {
let rdt_state = RdtState::new(store.clone());
start_change_forwarder(store, rdt_state.clients.clone());
Router::new()
.route("/rdt", get(websocket_handler))
.with_state(rdt_state)
}
pub fn router_with_rdt_state(rdt_state: RdtState) -> Router<RdtState> {
start_change_forwarder(rdt_state.store.clone(), rdt_state.clients.clone());
Router::new()
.route("/rdt", get(websocket_handler))
.with_state(rdt_state)
}
fn start_change_forwarder(
store: Arc<crate::document::DocumentStore>,
client_manager: Arc<ClientManager>,
) {
tokio::spawn(async move {
let mut change_rx = store.subscribe_to_changes();
tracing::info!("Started change forwarder for document store");
loop {
match change_rx.recv().await {
Ok((document_id, map_key, change_event)) => {
tracing::debug!(
"Forwarding change for document '{}', map '{}': {:?}",
document_id,
map_key,
change_event
);
let server_message = match change_event {
crate::protocol::ChangeEvent::Single(change) => {
crate::protocol::ServerMessage::MapChange {
document_id: document_id.clone(),
map_key: map_key.clone(),
change,
}
}
crate::protocol::ChangeEvent::Batch(changes) => {
crate::protocol::ServerMessage::BatchMapChange {
document_id: document_id.clone(),
map_key: map_key.clone(),
changes,
}
}
};
client_manager
.broadcast_to_subscribers(&document_id, &map_key, server_message)
.await;
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(skipped)) => {
tracing::warn!(
"Change forwarder lagged behind, skipped {} messages. Continuing...",
skipped
);
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
tracing::error!(
"Change forwarder channel closed - DocumentStore sender dropped"
);
break;
}
}
}
tracing::warn!("Change forwarder ended");
});
}
async fn websocket_handler(ws: WebSocketUpgrade, State(state): State<RdtState>) -> Response {
ws.on_upgrade(move |socket| WebSocketHandler::new(socket, state).handle())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::document::DocumentStore;
use serde_json::json;
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::time::{timeout, Duration};
#[tokio::test]
async fn test_change_forwarder_with_transactions() {
let store = Arc::new(DocumentStore::new());
let client_manager = Arc::new(ClientManager::new());
let (client_tx, mut client_rx) = mpsc::unbounded_channel();
let client_id = client_manager.register_client(client_tx).await;
client_manager
.subscribe_client(&client_id, "test-doc", "test-map")
.await;
start_change_forwarder(store.clone(), client_manager.clone());
tokio::time::sleep(Duration::from_millis(10)).await;
let doc = store.create_document("test-doc".to_string());
let map = doc.create_map("test-map".to_string());
map.insert("key1".to_string(), json!("value1"));
let message = timeout(Duration::from_millis(100), client_rx.recv())
.await
.expect("Should receive message")
.expect("Should have message");
match message {
crate::protocol::ServerMessage::MapChange {
document_id,
map_key,
change,
} => {
assert_eq!(document_id, "test-doc");
assert_eq!(map_key, "test-map");
assert!(matches!(change, crate::protocol::Change::Insert { .. }));
}
_ => panic!("Expected MapChange message, got: {:?}", message),
}
{
let transaction = map.start_transaction().unwrap();
map.insert("key2".to_string(), json!("value2"));
map.insert("key3".to_string(), json!("value3"));
map.remove("key2"); map.insert("key1".to_string(), json!("updated_value1"));
let result = timeout(Duration::from_millis(50), client_rx.recv()).await;
assert!(
result.is_err(),
"Should not receive message during transaction"
);
transaction.commit().unwrap();
}
let message = timeout(Duration::from_millis(100), client_rx.recv())
.await
.expect("Should receive message")
.expect("Should have message");
match message {
crate::protocol::ServerMessage::BatchMapChange {
document_id,
map_key,
changes,
} => {
assert_eq!(document_id, "test-doc");
assert_eq!(map_key, "test-map");
assert_eq!(changes.len(), 2);
let mut found_key3_insert = false;
let mut found_key1_update = false;
for change in changes {
match change {
crate::protocol::Change::Insert { key, value } if key == "key3" => {
assert_eq!(value, json!("value3"));
found_key3_insert = true;
}
crate::protocol::Change::Update {
key,
old_value,
new_value,
} if key == "key1" => {
assert_eq!(old_value, json!("value1"));
assert_eq!(new_value, json!("updated_value1"));
found_key1_update = true;
}
_ => panic!("Unexpected change: {:?}", change),
}
}
assert!(found_key3_insert, "Missing key3 insert");
assert!(found_key1_update, "Missing key1 update");
}
_ => panic!("Expected BatchMapChange message, got: {:?}", message),
}
}
}