use crate::{ auth, db, state::{AppState, NoteUpdate, PresenceUser, RoomEvent, SharedState}, }; use axum::{ extract::{ Path, State, WebSocketUpgrade, ws::{Message, WebSocket}, }, response::Response, }; use futures_util::{SinkExt, StreamExt}; use serde::{Deserialize, Serialize}; use std::time::{Duration, Instant}; use tracing::{debug, info, warn}; mod pad; pub use pad::upgrade_pad; #[derive(Debug, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] enum ClientMessage { Authenticate { password: Option, access_token: Option, nickname: Option, session_token: Option, guest_id: Option, color: Option, }, Update { content: String, owner_map: Option, }, Ping { nonce: u64, }, Chat { text: String, }, SetColor { color: Option, }, } #[derive(Debug, Serialize)] #[serde(tag = "type", rename_all = "snake_case")] enum ServerMessage { Authenticated { workspace_title: String, note_title: String, content: String, owner_map: String, access_level: String, }, Document { content: String, revision_id: i64, updated_at: String, author: Option, owner_map: String, }, Presence { users: Vec, }, Chat { sender: String, text: String, }, Pong { nonce: u64, }, Error { message: String, }, } // Merged from note.rs pub async fn upgrade( ws: WebSocketUpgrade, Path((workspace_slug, note_slug)): Path<(String, String)>, State(state): State, ) -> 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, ) { info!(%workspace_slug, %note_slug, "note websocket connected"); let Some(workspace) = db::find_workspace(&state.db, &workspace_slug) .await .ok() .flatten() else { warn!(%workspace_slug, %note_slug, "note websocket rejected: workspace not found"); let _ = send_error(&mut socket, "Workspace not found").await; return; }; let Some(note) = db::find_note(&state.db, workspace.id, ¬e_slug) .await .ok() .flatten() else { warn!(%workspace_slug, %note_slug, "note websocket rejected: note not found"); let _ = send_error(&mut socket, "Note not found").await; return; }; let (password, access_token, nickname, session_token, guest_id, color) = match socket.recv().await { Some(Ok(Message::Text(text))) => match serde_json::from_str::(&text) { Ok(ClientMessage::Authenticate { password, access_token, nickname, session_token, guest_id, color, }) => ( password, access_token, clean_nickname(nickname), session_token, clean_guest_id(guest_id), clean_color(color), ), _ => { let _ = send_error(&mut socket, "Wymagane uwierzytelnienie").await; return; } }, _ => return, }; let nickname = match auth::authorize_nickname(&state, nickname, session_token.clone()).await { Ok(value) => value, Err(message) => { let _ = send_error(&mut socket, &message).await; return; } }; let presence_identity = match session_token.as_deref() { Some(token) => auth::user_from_token(&state, token) .await .ok() .flatten() .map(|user| format!("user:{}", user.id)), None => guest_id.as_ref().and_then(|id| { nickname .as_ref() .map(|name| format!("guest:{id}:{}", name.to_lowercase())) }), }; let supplied_token = session_token.as_deref().or(access_token.as_deref()); let permission = auth::resource_permission(&state, "workspace", &workspace_slug, supplied_token) .await .ok() .flatten(); let anonymous_token_ok = permission.is_none() && crate::api::verify_resource_access_token( &state, "workspace", &workspace_slug, supplied_token, ) .await .unwrap_or(false); let password_ok = db::verify_workspace_password(&workspace, password.as_deref()); if workspace.is_private != 0 && permission.is_none() && !anonymous_token_ok { let _ = send_error(&mut socket, "Workspace not found").await; return; } if workspace.password_hash.is_some() && !password_ok && permission.is_none() && !anonymous_token_ok { warn!( workspace_id = workspace.id, note_id = note.id, "note websocket rejected: invalid workspace password" ); let _ = send_error(&mut socket, "Invalid password").await; return; } let write_allowed = permission.as_deref() == Some("rw") || anonymous_token_ok || password_ok || (workspace.is_private == 0 && workspace.password_hash.is_none() && permission.is_none()); info!(workspace_id = workspace.id, note_id = note.id, nickname = ?nickname, "note websocket authenticated"); if send( &mut socket, &ServerMessage::Authenticated { workspace_title: workspace.title.clone(), note_title: note.title.clone(), content: note.content.clone(), owner_map: note.owner_map.clone(), access_level: if write_allowed { "full".into() } else { "read_only".into() }, }, ) .await .is_err() { return; } let room_key = AppState::note_room_key(&workspace_slug, ¬e_slug); let channel = state.note_channel(&workspace_slug, ¬e_slug).await; let mut updates = channel.subscribe(); let display_name = nickname.clone().unwrap_or_else(|| "Guest".into()); let (connection_id, users) = state .join_room(&room_key, display_name.clone(), color, presence_identity) .await; let _ = channel.send(RoomEvent::Presence(users)); let mut last_chat = Instant::now() - Duration::from_secs(1); 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::(&text) { Ok(ClientMessage::Update{content,owner_map})=>{ if !write_allowed { let _=send_split(&mut sender,&ServerMessage::Error{message:"Read-only access".into()}).await; continue; } if content.len()>2_000_000 { let _=send_split(&mut sender,&ServerMessage::Error{message:"The document is too large".into()}).await; continue; } let owner_map=owner_map.unwrap_or_else(||"[]".into()); match db::save_revision(&state.db,note.id,workspace.id,&content,nickname.as_deref(),&owner_map).await { Ok((revision_id,updated_at))=>{let _=channel.send(RoomEvent::Document(NoteUpdate{content,revision_id,updated_at,author:nickname.clone(),owner_map}));} Err(error)=>warn!(%error, workspace_id = workspace.id, note_id = note.id, "failed to save revision"), } } Ok(ClientMessage::Ping{nonce})=>{ let _=send_split(&mut sender,&ServerMessage::Pong{nonce}).await; }, Ok(ClientMessage::Chat{text})=>{ let text=clean_chat(text); if !text.is_empty() && last_chat.elapsed() >= Duration::from_millis(500) { last_chat=Instant::now(); let _=channel.send(RoomEvent::Chat{sender:display_name.clone(),text}); } } Ok(ClientMessage::SetColor{color})=>{ let users=state.update_room_color(&room_key,connection_id,clean_color(color)).await; let _=channel.send(RoomEvent::Presence(users)); }, 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(RoomEvent::Document(update))=>if send_split(&mut sender,&ServerMessage::Document{content:update.content,revision_id:update.revision_id,updated_at:update.updated_at,author:update.author,owner_map:update.owner_map}).await.is_err(){break;}, Ok(RoomEvent::Presence(users))=>if send_split(&mut sender,&ServerMessage::Presence{users}).await.is_err(){break;}, Ok(RoomEvent::Chat{sender:chat_sender,text})=>if send_split(&mut sender,&ServerMessage::Chat{sender:chat_sender,text}).await.is_err(){break;}, Err(tokio::sync::broadcast::error::RecvError::Lagged(_))=>if let Ok(Some(current))=db::find_note(&state.db,workspace.id,¬e_slug).await { if send_split(&mut sender,&ServerMessage::Document{content:current.content,revision_id:0,updated_at:current.updated_at,author:None,owner_map:current.owner_map}).await.is_err(){break;} }, Err(tokio::sync::broadcast::error::RecvError::Closed)=>break, } } } let users = state.leave_room(&room_key, connection_id).await; let _ = channel.send(RoomEvent::Presence(users)); info!( workspace_id = workspace.id, note_id = note.id, "note websocket disconnected" ); } fn clean_nickname(value: Option) -> Option { value .map(|v| v.trim().chars().take(40).collect::()) .filter(|v| !v.is_empty()) } fn clean_guest_id(value: Option) -> Option { value .map(|v| v.trim().chars().take(64).collect::()) .filter(|v| { v.len() >= 16 && v.chars() .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_') }) } fn clean_color(value: Option) -> Option { value.map(|v| v.trim().to_ascii_lowercase()).filter(|v| { v.len() == 7 && v.starts_with('#') && v[1..].chars().all(|c| c.is_ascii_hexdigit()) }) } fn clean_chat(value: String) -> String { value .chars() .map(|c| { if matches!(c, '\r' | '\n' | '\0') { ' ' } else { c } }) .collect::() .trim() .chars() .take(1000) .collect() } 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> { socket .send(Message::Text( serde_json::to_string(message).unwrap().into(), )) .await } async fn send_split( sender: &mut futures_util::stream::SplitSink, message: &ServerMessage, ) -> Result<(), axum::Error> { sender .send(Message::Text( serde_json::to_string(message).unwrap().into(), )) .await } // Merged from pad.rs