Improve diff

This commit is contained in:
Andras Schmelczer 2026-05-09 16:27:48 +01:00
commit e5373ab2bb
23 changed files with 312 additions and 220 deletions

View file

@ -15,6 +15,7 @@ use super::{
};
use crate::{
app_state::websocket::models::DocumentWithCursors, config::database_config::DatabaseConfig,
errors::SyncServerError,
};
#[derive(Clone, Debug)]
@ -39,7 +40,7 @@ impl Cursors {
user_name: String,
device_id: &DeviceId,
document_to_cursors: Vec<DocumentWithCursors>,
) {
) -> Result<(), SyncServerError> {
let mut vault_to_cursors = self.vault_to_cursors.lock().await;
let all_device_cursors = vault_to_cursors
@ -54,7 +55,7 @@ impl Cursors {
}));
drop(vault_to_cursors); // Explicitly drop the lock before broadcasting to avoid deadlock
self.broadcast_cursors_for_vault(&vault_id).await;
self.broadcast_cursors_for_vault(&vault_id).await
}
pub async fn get_cursors(&self, vault_id: &VaultId) -> Vec<ClientCursors> {
@ -76,15 +77,17 @@ impl Cursors {
loop {
tokio::select! {
() = tokio::time::sleep(Duration::from_secs(1)) => {
self.remove_expired_cursors().await;
self.remove_expired_cursors().await?;
}
Ok(()) = shutdown.changed() => break,
}
}
Ok::<(), SyncServerError>(())
});
}
async fn remove_expired_cursors(&self) {
async fn remove_expired_cursors(&self) -> Result<(), SyncServerError> {
let changed_vaults: Vec<VaultId> = {
let mut vault_to_cursors = self.vault_to_cursors.lock().await;
@ -104,11 +107,13 @@ impl Cursors {
};
for vault_id in &changed_vaults {
self.broadcast_cursors_for_vault(vault_id).await;
self.broadcast_cursors_for_vault(vault_id).await?;
}
Ok(())
}
async fn broadcast_cursors_for_vault(&self, vault_id: &VaultId) {
async fn broadcast_cursors_for_vault(&self, vault_id: &VaultId) -> Result<(), SyncServerError> {
let client_cursors: Vec<ClientCursors> = {
let vault_to_cursors = self.vault_to_cursors.lock().await;
vault_to_cursors
@ -124,10 +129,14 @@ impl Cursors {
clients: client_cursors,
},
)),
);
)
}
pub async fn remove_cursors_of_device(&self, vault_id: &VaultId, device_id: &DeviceId) {
pub async fn remove_cursors_of_device(
&self,
vault_id: &VaultId,
device_id: &DeviceId,
) -> Result<(), SyncServerError> {
let changed = {
let mut vault_to_cursors = self.vault_to_cursors.lock().await;
@ -145,8 +154,9 @@ impl Cursors {
};
if changed {
self.broadcast_cursors_for_vault(vault_id).await;
self.broadcast_cursors_for_vault(vault_id).await?;
}
Ok(())
}
}

View file

@ -5,7 +5,7 @@ use std::{
sync::atomic::{AtomicU64, Ordering},
};
use anyhow::{Context as _, Result};
use anyhow::{Context as _, Result, anyhow};
use log::info;
use models::{
DocumentId, DocumentVersionWithoutContent, StoredDocumentVersion, VaultId, VaultUpdateId,
@ -132,6 +132,12 @@ impl WriteTransaction {
}
Ok(())
}
pub fn connection_mut(&mut self) -> Result<&mut SqliteConnection> {
self.conn
.as_deref_mut()
.context("WriteTransaction already consumed")
}
}
impl Drop for WriteTransaction {
@ -147,25 +153,6 @@ impl Drop for WriteTransaction {
}
}
impl std::ops::Deref for WriteTransaction {
type Target = SqliteConnection;
fn deref(&self) -> &Self::Target {
self.conn
.as_ref()
.expect("BUG: WriteTransaction dereferenced after being consumed")
.deref()
}
}
impl std::ops::DerefMut for WriteTransaction {
fn deref_mut(&mut self) -> &mut Self::Target {
self.conn
.as_mut()
.expect("BUG: WriteTransaction dereferenced after being consumed")
.deref_mut()
}
}
/// Ensure the connection has no leftover open transaction (e.g. from a
/// `WriteTransaction` that was dropped without commit/rollback). ROLLBACK
/// is a harmless no-op if no transaction is active.
@ -797,7 +784,7 @@ impl Database {
let _send_guard = self.broadcasts.acquire_send_lock(vault_id).await;
query
.execute(&mut *transaction)
.execute(transaction.connection_mut()?)
.await
.context("Cannot insert document version")?;
@ -821,7 +808,8 @@ impl Database {
} else {
WebSocketServerMessageWithOrigin::with_origin(version.device_id.clone(), envelope)
};
self.broadcasts.send_document_update(vault_id, with_origin);
self.broadcasts
.send_document_update(vault_id, with_origin)?;
Ok(())
}

View file

@ -7,7 +7,11 @@ use log::{debug, info, warn};
use tokio::sync::{Mutex, broadcast};
use super::models::{WebSocketServerMessage, WebSocketServerMessageWithOrigin};
use crate::{app_state::database::models::VaultId, config::server_config::ServerConfig};
use crate::{
app_state::database::models::VaultId,
config::server_config::ServerConfig,
errors::{SyncServerError, client_error, server_error},
};
#[derive(Debug, Clone)]
pub struct Broadcasts {
@ -60,30 +64,31 @@ impl Broadcasts {
pub fn get_receiver(
&self,
vault: VaultId,
vault: &VaultId,
max_clients: usize,
) -> Result<broadcast::Receiver<WebSocketServerMessageWithOrigin>, crate::errors::SyncServerError>
{
) -> Result<broadcast::Receiver<WebSocketServerMessageWithOrigin>, SyncServerError> {
let mut tx_map = self
.tx
.lock()
.expect("broadcasts.tx mutex poisoned — a previous holder panicked");
.map_err(|_| server_error(anyhow::anyhow!("broadcasts.tx mutex poisoned")))?;
let count_before_prune = tx_map
.get(&vault)
.get(vault)
.map_or(0, tokio::sync::broadcast::Sender::receiver_count);
let pruned = Self::prune_inactive_vaults(&mut tx_map);
let pruned_self = pruned.contains(&vault);
let pruned_self = pruned
.iter()
.any(|pruned_vault| pruned_vault.as_str() == vault);
let sender = tx_map
.entry(vault.clone())
.entry(vault.to_owned())
.or_insert_with(|| broadcast::channel(self.broadcast_channel_capacity).0);
// Hold the lock across the count check *and* the subscribe so the
// `max_clients` cap is atomic: two concurrent callers can't both
// observe `receiver_count() < max_clients` and both subscribe.
if sender.receiver_count() >= max_clients {
return Err(crate::errors::client_error(anyhow::anyhow!(
return Err(client_error(anyhow::anyhow!(
"Vault has reached the maximum number of clients ({max_clients})"
)));
}
@ -100,8 +105,13 @@ impl Broadcasts {
/// Notify all clients (who are subscribed to the vault) about an update.
/// Synchronous: safe to invoke from a handler between `commit()` and
/// function return without worrying about task cancellation dropping
/// the broadcast mid-flight. Failures are logged, never propagated.
pub fn send_document_update(&self, vault: VaultId, document: WebSocketServerMessageWithOrigin) {
/// the broadcast mid-flight. Mutex poison is returned; send failures
/// are logged because they can happen when receivers disconnect.
pub fn send_document_update(
&self,
vault: &str,
document: WebSocketServerMessageWithOrigin,
) -> Result<(), SyncServerError> {
let vault_update_id = match &document.message {
WebSocketServerMessage::VaultUpdate(u) => Some(u.document.vault_update_id),
WebSocketServerMessage::CursorPositions(_) => None,
@ -110,18 +120,21 @@ impl Broadcasts {
WebSocketServerMessage::VaultUpdate(u) => Some(u.document.is_deleted),
WebSocketServerMessage::CursorPositions(_) => None,
};
let mut tx_map = self
.tx
.lock()
.expect("broadcasts.tx mutex poisoned — a previous holder panicked");
let mut tx_map = self.tx.lock().map_err(|_| {
server_error(anyhow::anyhow!(
"broadcasts.tx mutex poisoned; skipping document update broadcast"
))
})?;
let count_before_prune = tx_map
.get(&vault)
.get(vault)
.map_or(0, tokio::sync::broadcast::Sender::receiver_count);
let pruned = Self::prune_inactive_vaults(&mut tx_map);
let pruned_self = pruned.contains(&vault);
let pruned_self = pruned
.iter()
.any(|pruned_vault| pruned_vault.as_str() == vault);
let sender = tx_map
.entry(vault.clone())
.entry(vault.to_owned())
.or_insert_with(|| broadcast::channel(self.broadcast_channel_capacity).0);
let count_before_send = sender.receiver_count();
@ -131,7 +144,7 @@ impl Broadcasts {
"[BCAST] send_document_update vault={vault} vuid={vault_update_id:?} is_deleted={is_deleted:?} count_before_prune={count_before_prune} pruned_self={pruned_self} count_before_send=0 SKIPPED"
);
debug!("Skipping broadcast, no clients connected for vault `{vault}`");
return;
return Ok(());
}
let send_result = sender.send(document);
@ -143,5 +156,6 @@ impl Broadcasts {
"[BCAST] send_document_update vault={vault} vuid={vault_update_id:?} is_deleted={is_deleted:?} count_before_prune={count_before_prune} pruned_self={pruned_self} count_before_send={count_before_send} FAILED err={e}"
),
}
Ok(())
}
}