first commit

This commit is contained in:
Mateusz Gruszczyński
2026-07-17 15:29:08 +02:00
commit 771494671b
35 changed files with 5020 additions and 0 deletions
+292
View File
@@ -0,0 +1,292 @@
use axum::{
extract::{
ws::{Message, WebSocket},
Path, State, WebSocketUpgrade,
},
response::Response,
};
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
use tracing::{debug, warn};
use crate::{
db,
state::{NoteUpdate, SharedState},
};
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ClientMessage {
Authenticate { password: Option<String> },
Update { content: String },
}
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ServerMessage {
Authenticated {
workspace_title: String,
note_title: String,
content: String,
},
Document {
content: String,
revision_id: i64,
updated_at: String,
},
Error { message: String },
}
pub async fn upgrade(
ws: WebSocketUpgrade,
Path((workspace_slug, note_slug)): Path<(String, String)>,
State(state): State<SharedState>,
) -> Response {
ws.on_upgrade(move |socket| handle_socket(socket, state, workspace_slug, note_slug))
}
async fn handle_socket(
mut socket: WebSocket,
state: SharedState,
workspace_slug: String,
note_slug: String,
) {
let Some(workspace) = db::find_workspace(&state.db, &workspace_slug)
.await
.ok()
.flatten()
else {
let _ = send_error(&mut socket, "Nie znaleziono workspace").await;
return;
};
let Some(note) = db::find_note(&state.db, workspace.id, &note_slug)
.await
.ok()
.flatten()
else {
let _ = send_error(&mut socket, "Nie znaleziono notatki").await;
return;
};
let password = match socket.recv().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Authenticate { password }) => password,
_ => {
let _ = send_error(&mut socket, "Wymagane uwierzytelnienie").await;
return;
}
},
_ => return,
};
if !db::verify_workspace_password(&workspace, password.as_deref()) {
let _ = send_error(&mut socket, "Nieprawidłowe hasło").await;
return;
}
if send(
&mut socket,
&ServerMessage::Authenticated {
workspace_title: workspace.title.clone(),
note_title: note.title.clone(),
content: note.content.clone(),
},
)
.await
.is_err()
{
return;
}
let channel = state.note_channel(&workspace_slug, &note_slug).await;
let mut updates = channel.subscribe();
let (mut sender, mut receiver) = socket.split();
loop {
tokio::select! {
incoming = receiver.next() => {
match incoming {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Update { content }) => {
if content.len() > 2_000_000 {
let _ = send_split(&mut sender, &ServerMessage::Error { message: "Dokument jest zbyt duży".into() }).await;
continue;
}
match db::save_revision(&state.db, note.id, workspace.id, &content).await {
Ok((revision_id, updated_at)) => {
let _ = channel.send(NoteUpdate { content, revision_id, updated_at });
}
Err(error) => warn!(%error, "failed to save revision"),
}
}
Ok(ClientMessage::Authenticate { .. }) => {}
Err(error) => warn!(%error, "invalid WebSocket message"),
},
Some(Ok(Message::Close(_))) | None => break,
Some(Ok(_)) => {}
Some(Err(error)) => {
debug!(%error, "WebSocket receive error");
break;
}
}
}
update = updates.recv() => match update {
Ok(update) => {
if send_split(&mut sender, &ServerMessage::Document {
content: update.content,
revision_id: update.revision_id,
updated_at: update.updated_at,
}).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
if let Ok(Some(current)) = db::find_note(&state.db, workspace.id, &note_slug).await {
if send_split(&mut sender, &ServerMessage::Document {
content: current.content,
revision_id: 0,
updated_at: current.updated_at,
}).await.is_err() {
break;
}
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
}
}
}
}
async fn send_error(socket: &mut WebSocket, message: &str) -> Result<(), axum::Error> {
send(socket, &ServerMessage::Error { message: message.into() }).await
}
async fn send(socket: &mut WebSocket, message: &ServerMessage) -> Result<(), axum::Error> {
let payload = serde_json::to_string(message).expect("serializing server message cannot fail");
socket.send(Message::Text(payload.into())).await
}
async fn send_split(
sender: &mut futures_util::stream::SplitSink<WebSocket, Message>,
message: &ServerMessage,
) -> Result<(), axum::Error> {
let payload = serde_json::to_string(message).expect("serializing server message cannot fail");
sender.send(Message::Text(payload.into())).await
}
#[derive(Debug, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum PadServerMessage {
Authenticated { title: String, content: String },
Document { content: String, revision_id: i64, updated_at: String },
Error { message: String },
}
pub async fn upgrade_pad(
ws: WebSocketUpgrade,
Path(slug): Path<String>,
State(state): State<SharedState>,
) -> Response {
ws.on_upgrade(move |socket| handle_pad_socket(socket, state, slug))
}
async fn handle_pad_socket(mut socket: WebSocket, state: SharedState, slug: String) {
let Some(pad) = db::find_pad(&state.db, &slug).await.ok().flatten() else {
let _ = send_pad(&mut socket, &PadServerMessage::Error { message: "Nie znaleziono notatki".into() }).await;
return;
};
let password = match socket.recv().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Authenticate { password }) => password,
_ => {
let _ = send_pad(&mut socket, &PadServerMessage::Error { message: "Wymagane uwierzytelnienie".into() }).await;
return;
}
},
_ => return,
};
if !db::verify_pad_password(&pad, password.as_deref()) {
let _ = send_pad(&mut socket, &PadServerMessage::Error { message: "Nieprawidłowe hasło".into() }).await;
return;
}
if send_pad(&mut socket, &PadServerMessage::Authenticated {
title: pad.title.clone(),
content: pad.content.clone(),
}).await.is_err() {
return;
}
let channel = state.pad_channel(&slug).await;
let mut updates = channel.subscribe();
let (mut sender, mut receiver) = socket.split();
loop {
tokio::select! {
incoming = receiver.next() => {
match incoming {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Update { content }) => {
if content.len() > 2_000_000 {
let _ = send_pad_split(&mut sender, &PadServerMessage::Error { message: "Dokument jest zbyt duży".into() }).await;
continue;
}
match db::save_pad_revision(&state.db, pad.id, &content).await {
Ok((revision_id, updated_at)) => {
let _ = channel.send(NoteUpdate { content, revision_id, updated_at });
}
Err(error) => warn!(%error, "failed to save pad revision"),
}
}
Ok(ClientMessage::Authenticate { .. }) => {}
Err(error) => warn!(%error, "invalid pad WebSocket message"),
},
Some(Ok(Message::Close(_))) | None => break,
Some(Ok(_)) => {}
Some(Err(error)) => {
debug!(%error, "pad WebSocket receive error");
break;
}
}
}
update = updates.recv() => match update {
Ok(update) => {
if send_pad_split(&mut sender, &PadServerMessage::Document {
content: update.content,
revision_id: update.revision_id,
updated_at: update.updated_at,
}).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {
if let Ok(Some(current)) = db::find_pad(&state.db, &slug).await {
if send_pad_split(&mut sender, &PadServerMessage::Document {
content: current.content,
revision_id: 0,
updated_at: current.updated_at,
}).await.is_err() {
break;
}
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
}
}
}
}
async fn send_pad(socket: &mut WebSocket, message: &PadServerMessage) -> Result<(), axum::Error> {
let payload = serde_json::to_string(message).expect("serializing pad server message cannot fail");
socket.send(Message::Text(payload.into())).await
}
async fn send_pad_split(
sender: &mut futures_util::stream::SplitSink<WebSocket, Message>,
message: &PadServerMessage,
) -> Result<(), axum::Error> {
let payload = serde_json::to_string(message).expect("serializing pad server message cannot fail");
sender.send(Message::Text(payload.into())).await
}