Implement cursor broadcasting backend

This commit is contained in:
Andras Schmelczer 2025-06-01 09:50:52 +01:00
parent 483e03e2de
commit eb1cc61042
No known key found for this signature in database
GPG key ID: FC8F2C3D3D1A718C
19 changed files with 488 additions and 191 deletions

View file

@ -6,64 +6,52 @@ use axum::{
},
response::Response,
};
use futures::{
sink::SinkExt,
stream::{SplitSink, StreamExt},
};
use futures::stream::StreamExt;
use log::{error, info, warn};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use serde::Deserialize;
use super::auth::auth;
use crate::{
app_state::{
AppState,
database::models::{DeviceId, DocumentVersionWithoutContent, VaultId, VaultUpdateId},
database::models::VaultId,
websocket::{
models::{
CursorPositionFromServer, WebSocketClientMessage, WebSocketServerMessage,
WebSocketVaultUpdate,
},
utils::{get_handshake, get_unseen_documents, send_update_over_websocket},
},
},
errors::{SyncServerError, server_error, unauthenticated_error},
errors::{SyncServerError, client_error, server_error, unauthenticated_error},
utils::normalize::normalize,
};
// This is required for aide to infer the path parameter types and names
#[derive(Deserialize, JsonSchema)]
pub struct WebsocketPathParams {
pub struct WebSocketPathParams {
#[serde(deserialize_with = "normalize")]
vault_id: VaultId,
}
pub async fn websocket_handler(
ws: WebSocketUpgrade,
Path(WebsocketPathParams { vault_id }): Path<WebsocketPathParams>,
Path(WebSocketPathParams { vault_id }): Path<WebSocketPathParams>,
State(state): State<AppState>,
) -> Result<Response, SyncServerError> {
Ok(ws.on_upgrade(move |socket| websocket_wrapped(state, socket, vault_id)))
}
async fn websocket_wrapped(state: AppState, stream: WebSocket, vault_id: VaultId) {
info!("Websocket connection opened on vault '{vault_id}'");
info!("WebSocket connection opened on vault '{vault_id}'");
let result = websocket(state, stream, vault_id.clone()).await;
if let Err(err) = result {
error!("Websocket connection error on vault '{vault_id}': {err}");
error!("WebSocket connection error on vault '{vault_id}': {err}");
}
warn!("Websocket connection closed on vault '{vault_id}'");
}
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct WebsocketHandshake {
pub token: String,
pub device_id: DeviceId,
pub last_seen_vault_update_id: Option<VaultUpdateId>,
}
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
struct WebsocketVaultUpdate {
pub documents: Vec<DocumentVersionWithoutContent>,
pub is_initial_sync: bool,
warn!("WebSocket connection closed on vault '{vault_id}'");
}
async fn websocket(
@ -73,68 +61,71 @@ async fn websocket(
) -> Result<(), SyncServerError> {
let (mut sender, mut receiver) = stream.split();
let handshake = if let Some(Ok(Message::Text(token))) = receiver.next().await {
let handshake: WebsocketHandshake = serde_json::from_str(&token)
.context("Failed to parse token")
.map_err(server_error)?;
auth(&state, handshake.token.trim(), &vault_id)?;
handshake
let handshake = if let Some(Ok(message)) = receiver.next().await {
get_handshake(&state, &vault_id, message)?
} else {
return Err(unauthenticated_error(anyhow::anyhow!(
"Failed to authenticate"
"Failed to authenticate due to invalid message"
)));
};
let mut rx = state.broadcasts.get_receiver(vault_id.clone()).await;
let documents = if let Some(update_id) = handshake.last_seen_vault_update_id {
state
.database
.get_latest_documents_since(&vault_id, update_id, None)
.await
.map_err(server_error)
} else {
state
.database
.get_latest_documents(&vault_id, None)
.await
.map_err(server_error)
}?;
send_update_over_websocket(
&WebsocketVaultUpdate {
documents,
&WebSocketServerMessage::VaultUpdate(WebSocketVaultUpdate {
documents: get_unseen_documents(&state, &vault_id, handshake.last_seen_vault_update_id)
.await?,
is_initial_sync: true,
},
}),
&mut sender,
)
.await?;
send_update_over_websocket(
&WebSocketServerMessage::CursorPositions(CursorPositionFromServer {
clients: state.cursors.get_cursors(&vault_id).await,
}),
&mut sender,
)
.await?;
let device_id = handshake.device_id.clone();
let mut send_task = tokio::spawn(async move {
while let Ok(update) = rx.recv().await {
if Some(&handshake.device_id) == update.origin_device_id.as_ref() {
if Some(&device_id) == update.origin_device_id.as_ref() {
continue;
}
send_update_over_websocket(
&WebsocketVaultUpdate {
documents: vec![update.document],
is_initial_sync: false,
},
&mut sender,
)
.await?;
send_update_over_websocket(&update.message, &mut sender).await?;
}
Ok::<(), SyncServerError>(())
});
let mut recv_task =
tokio::spawn(
async move { while let Some(Ok(Message::Text(_text))) = receiver.next().await {} },
);
let device_id = handshake.device_id.clone();
let mut recv_task = tokio::spawn(async move {
while let Some(Ok(Message::Text(message))) = receiver.next().await {
let message: WebSocketClientMessage = serde_json::from_str(&message)
.context("Failed to parse message")
.map_err(server_error)?;
match message {
WebSocketClientMessage::Handshake(_) => {
return Err(client_error(anyhow::anyhow!(
"Unexpected handshake message"
)));
}
WebSocketClientMessage::CursorPositions(cursors) => {
state
.cursors
.update_cursors(vault_id.clone(), &device_id, cursors.document_to_cursors)
.await;
}
}
}
Ok::<(), SyncServerError>(())
});
tokio::select! {
_ = &mut send_task => recv_task.abort(),
@ -143,28 +134,13 @@ async fn websocket(
send_task
.await
.context("Websocket send task failed")
.context("WebSocket send task failed")
.map_err(server_error)??;
recv_task
.await
.context("Websocket receive task failed")
.map_err(server_error)?;
.context("WebSocket receive task failed")
.map_err(server_error)??;
Ok(())
}
async fn send_update_over_websocket(
update: &WebsocketVaultUpdate,
sender: &mut SplitSink<WebSocket, Message>,
) -> Result<(), SyncServerError> {
let serialized_update = serde_json::to_string(update)
.context("Failed to serialize update")
.map_err(server_error)?;
sender
.send(Message::Text(serialized_update))
.await
.context("Failed to send message over websocket")
.map_err(server_error)
}