From 917be995a48e4c0ef34c853f6b8e3cf7e4a96a7a Mon Sep 17 00:00:00 2001 From: Blomios Date: Wed, 15 Jul 2026 18:09:14 +0200 Subject: [PATCH] =?UTF-8?q?feat(server):=20endpoint=20PTY=20WebSocket=20au?= =?UTF-8?q?thentifi=C3=A9=20sur=20idea=20--serve=20(#13)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Lot B5 du chantier server/client mode : le serveur --serve expose le terminal distant via un endpoint WebSocket, préalable au client xterm (F3). - Handshake WebSocket RFC 6455 fait main, sans nouvelle dépendance. - Auth de l'upgrade par cookie de session + Origin strict. - Frames PTY (open/attach/close/ping), réattache, scrollback, backpressure. - 37 tests server dont les refus de sécurité et les handlers open/attach/close/ping. Clippy propre (result_large_err corrigé). Validé : app-tauri 277 tests verts, backend 28, cohérence des frames B5↔F3 confirmée, desktop non régressé. Réserve connue : le round-trip socket réel relève d'une validation live hors sandbox. Co-Authored-By: Claude Opus 4.8 --- crates/app-tauri/src/server.rs | 1260 +++++++++++++++++++++++++++++++- 1 file changed, 1254 insertions(+), 6 deletions(-) diff --git a/crates/app-tauri/src/server.rs b/crates/app-tauri/src/server.rs index c9953cd..0e78821 100644 --- a/crates/app-tauri/src/server.rs +++ b/crates/app-tauri/src/server.rs @@ -9,8 +9,10 @@ use std::env; use std::net::SocketAddr; use std::path::PathBuf; use std::process::ExitCode; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; +use backend::stream::{OutputBridge, OutputSink, OutputSinkError}; use bytes::Bytes; use cookie::{Cookie, SameSite}; use http::header::{HeaderValue, CONTENT_TYPE, COOKIE, ORIGIN, SET_COOKIE}; @@ -18,23 +20,33 @@ use http::header::{HeaderValue, CONTENT_TYPE, COOKIE, ORIGIN, SET_COOKIE}; use http::Request; use http::{HeaderMap, Method, Response, StatusCode, Uri}; use http_body_util::{BodyExt, Full}; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpListener; +use tokio::sync::mpsc; use uuid::Uuid; -use application::{GetProjectWorkStateInput, OpenProjectInput}; -use domain::Project; +use application::{ + CloseTerminalInput, GetProjectWorkStateInput, OpenProjectInput, ResizeTerminalInput, + WriteToTerminalInput, +}; +use domain::ports::PtyHandle; +use domain::{Project, SessionId}; use crate::dto::{ - parse_project_id, ErrorDto, HealthRequestDto, HealthResponseDto, ProjectDto, ProjectListDto, - ProjectWorkStateDto, + parse_project_id, parse_session_id, ErrorDto, HealthRequestDto, HealthResponseDto, + OpenTerminalRequestDto, ProjectDto, ProjectListDto, ProjectWorkStateDto, TerminalSessionDto, }; +use crate::pty::PtyChunk; use crate::state::AppState; const DEFAULT_LISTEN: &str = "127.0.0.1:17373"; const SESSION_COOKIE: &str = "idea_session"; +const WS_PATH: &str = "/api/ws"; +const WS_MAGIC: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; +const WS_MAX_PAYLOAD: usize = 64 * 1024; +const WS_OUTPUT_BUFFER: usize = 512; type ResponseBody = Full; @@ -176,6 +188,7 @@ struct ServerState { app: AppState, pairing_code: String, sessions: Mutex>, + ws_pty_bridge: Arc>, } impl ServerState { @@ -185,6 +198,7 @@ impl ServerState { config, pairing_code: new_pairing_code(), sessions: Mutex::new(HashSet::new()), + ws_pty_bridge: Arc::new(OutputBridge::new()), } } @@ -195,6 +209,7 @@ impl ServerState { config, pairing_code: pairing_code.into(), sessions: Mutex::new(HashSet::new()), + ws_pty_bridge: Arc::new(OutputBridge::new()), } } @@ -261,6 +276,10 @@ async fn handle_tcp_connection( }; let (method, uri, headers, content_length) = parse_http_request_head(&buffer[..header_end])?; + if method == Method::GET && uri.path() == WS_PATH { + return handle_ws_upgrade(stream, headers, state).await; + } + let body_start = header_end + 4; while buffer.len() < body_start + content_length { let read = stream @@ -395,6 +414,373 @@ async fn write_http_response( .map_err(|err| format!("failed to write response: {err}")) } +async fn handle_ws_upgrade( + mut stream: tokio::net::TcpStream, + headers: HeaderMap, + state: Arc, +) -> Result<(), String> { + let accept = match validate_ws_upgrade(&headers, &state) { + Ok(accept) => accept, + Err(response) => return write_http_response(&mut stream, *response).await, + }; + + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\n\ + upgrade: websocket\r\n\ + connection: Upgrade\r\n\ + sec-websocket-accept: {accept}\r\n\r\n" + ); + stream + .write_all(response.as_bytes()) + .await + .map_err(|err| format!("failed to write websocket upgrade: {err}"))?; + run_ws_connection(stream, state).await +} + +fn validate_ws_upgrade( + headers: &HeaderMap, + state: &ServerState, +) -> Result>> { + let origin = validate_request_origin(headers, &state.config)?; + let Some(token) = session_cookie(headers) else { + return Err(Box::new(error_response( + StatusCode::UNAUTHORIZED, + "UNAUTHORIZED", + "missing session cookie", + origin.as_deref(), + ))); + }; + if !state.has_session(&token) { + return Err(Box::new(error_response( + StatusCode::UNAUTHORIZED, + "UNAUTHORIZED", + "invalid session cookie", + origin.as_deref(), + ))); + } + if !header_contains_token(headers, "upgrade", "websocket") + || !header_contains_token(headers, "connection", "upgrade") + { + return Err(Box::new(error_response( + StatusCode::BAD_REQUEST, + "INVALID", + "missing websocket upgrade headers", + origin.as_deref(), + ))); + } + let Some(key) = headers + .get("sec-websocket-key") + .and_then(|value| value.to_str().ok()) + else { + return Err(Box::new(error_response( + StatusCode::BAD_REQUEST, + "INVALID", + "missing Sec-WebSocket-Key", + origin.as_deref(), + ))); + }; + Ok(websocket_accept(key)) +} + +fn header_contains_token(headers: &HeaderMap, name: &str, needle: &str) -> bool { + headers + .get(name) + .and_then(|value| value.to_str().ok()) + .is_some_and(|value| { + value + .split(',') + .any(|token| token.trim().eq_ignore_ascii_case(needle)) + }) +} + +async fn run_ws_connection( + stream: tokio::net::TcpStream, + state: Arc, +) -> Result<(), String> { + let (mut reader, mut writer) = stream.into_split(); + let (tx, mut rx) = mpsc::channel::(WS_OUTPUT_BUFFER); + let owned = Arc::new(Mutex::new(Vec::<(SessionId, u64)>::new())); + let writer_task = tokio::spawn(async move { + while let Some(frame) = rx.recv().await { + let text = serde_json::to_vec(&frame) + .map_err(|err| format!("failed to encode server frame: {err}"))?; + let bytes = encode_ws_frame(WsOpcode::Text, &text); + writer + .write_all(&bytes) + .await + .map_err(|err| format!("failed to write websocket frame: {err}"))?; + } + Ok::<(), String>(()) + }); + + loop { + let frame = match read_ws_frame(&mut reader).await { + Ok(frame) => frame, + Err(err) => { + let _ = tx + .send(ServerFrame::error( + None, + None, + "WS_PROTOCOL", + err.to_string(), + )) + .await; + break; + } + }; + match frame.opcode { + WsOpcode::Text => { + let parsed = serde_json::from_slice::(&frame.payload) + .map_err(|err| format!("invalid client frame JSON: {err}")); + match parsed { + Ok(frame) => { + handle_client_frame(frame, &state, &tx, &owned).await; + } + Err(err) => { + let _ = tx + .send(ServerFrame::error(None, None, "INVALID", err)) + .await; + } + } + } + WsOpcode::Ping => { + let _ = tx.send(ServerFrame::pong()).await; + } + WsOpcode::Close => break, + WsOpcode::Pong => {} + WsOpcode::Binary => { + let _ = tx + .send(ServerFrame::error( + None, + None, + "UNSUPPORTED", + "binary websocket frames are not supported", + )) + .await; + } + } + } + + if let Ok(owned) = owned.lock() { + for (session, gen) in owned.iter() { + state.ws_pty_bridge.unregister_if(session, *gen); + } + } + drop(tx); + writer_task + .await + .map_err(|err| format!("websocket writer task failed: {err}"))? +} + +async fn handle_client_frame( + frame: ClientFrame, + state: &Arc, + tx: &mpsc::Sender, + owned: &Arc>>, +) { + let result = match frame.kind.as_str() { + "terminal.open" | "open_terminal" => ws_open_terminal(&frame, state, tx, owned).await, + "terminal.attach" | "attach_terminal" => ws_attach_terminal(&frame, state, tx, owned).await, + "terminal.input" | "input" => ws_input(&frame, state), + "terminal.resize" | "resize" => ws_resize(&frame, state), + "terminal.detach" | "detach" => ws_detach(&frame, state, owned), + "terminal.close" | "close" => match ws_close(&frame, state, owned).await { + Ok((session_id, exit_code)) => { + let _ = tx + .send(ServerFrame::status(session_id, "exited", exit_code)) + .await; + Ok(()) + } + Err(err) => Err(err), + }, + "ping" => { + let _ = tx.send(ServerFrame::pong()).await; + Ok(()) + } + _ => Err(ErrorDto { + code: "UNKNOWN_FRAME".to_owned(), + message: format!("unknown websocket frame kind: {}", frame.kind), + }), + }; + + if let Err(err) = result { + let _ = tx + .send(ServerFrame::error( + Some(frame.id), + payload_session_id(&frame.payload), + err.code, + err.message, + )) + .await; + } +} + +async fn ws_open_terminal( + frame: &ClientFrame, + state: &Arc, + tx: &mpsc::Sender, + owned: &Arc>>, +) -> Result<(), ErrorDto> { + let request_value = frame + .payload + .get("request") + .cloned() + .unwrap_or_else(|| frame.payload.clone()); + let request: OpenTerminalRequestDto = + serde_json::from_value(request_value).map_err(invalid_args_error)?; + let output = state + .app + .open_terminal + .execute(request.into()) + .await + .map_err(ErrorDto::from)?; + let dto = TerminalSessionDto::from(output); + let sid = parse_session_id(&dto.session_id)?; + let sink = WsPtySink::new(sid, tx.clone(), 0); + tx.send(ServerFrame::attached(&frame.id, &dto, Vec::new(), 0, false)) + .await + .map_err(|_| ErrorDto { + code: "WS_CLOSED".to_owned(), + message: "websocket output closed".to_owned(), + })?; + attach_sink_and_pump(state, sid, sink, owned) +} + +async fn ws_attach_terminal( + frame: &ClientFrame, + state: &Arc, + tx: &mpsc::Sender, + owned: &Arc>>, +) -> Result<(), ErrorDto> { + let session_id = payload_string(&frame.payload, "sessionId")?; + let sid = parse_session_id(&session_id)?; + let handle = PtyHandle { session_id: sid }; + let scrollback = state + .app + .pty_port + .scrollback(&handle) + .map_err(|err| ErrorDto::from(application::AppError::from(err)))?; + let rows = payload_u16(&frame.payload, "rows").unwrap_or(24); + let cols = payload_u16(&frame.payload, "cols").unwrap_or(80); + let dto = TerminalSessionDto { + session_id, + cwd: String::new(), + rows, + cols, + assigned_conversation_id: None, + engine_session_id: None, + cell_kind: crate::dto::CellKind::Pty, + }; + let next_seq = u64::from(!scrollback.is_empty()); + tx.send(ServerFrame::attached( + &frame.id, + &dto, + scrollback.clone(), + next_seq, + frame.payload.get("lastSeq").is_some(), + )) + .await + .map_err(|_| ErrorDto { + code: "WS_CLOSED".to_owned(), + message: "websocket output closed".to_owned(), + })?; + let sink = WsPtySink::new(sid, tx.clone(), next_seq); + attach_sink_and_pump(state, sid, sink, owned) +} + +fn attach_sink_and_pump( + state: &Arc, + sid: SessionId, + sink: WsPtySink, + owned: &Arc>>, +) -> Result<(), ErrorDto> { + let gen = state.ws_pty_bridge.register(sid, Arc::new(sink)); + if let Ok(mut owned) = owned.lock() { + owned.push((sid, gen)); + } + let handle = PtyHandle { session_id: sid }; + let stream = state + .app + .pty_port + .subscribe_output(&handle) + .map_err(|err| ErrorDto::from(application::AppError::from(err)))?; + let bridge = Arc::clone(&state.ws_pty_bridge); + std::thread::spawn(move || { + for chunk in stream { + if !bridge.send_output(&sid, chunk) { + break; + } + } + bridge.unregister_if(&sid, gen); + }); + Ok(()) +} + +fn ws_input(frame: &ClientFrame, state: &Arc) -> Result<(), ErrorDto> { + let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?; + let bytes = payload_bytes(&frame.payload)?; + state + .app + .write_terminal + .execute(WriteToTerminalInput { + session_id: sid, + data: bytes, + }) + .map_err(ErrorDto::from) +} + +fn ws_resize(frame: &ClientFrame, state: &Arc) -> Result<(), ErrorDto> { + let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?; + let rows = payload_u16(&frame.payload, "rows").ok_or_else(|| ErrorDto { + code: "INVALID".to_owned(), + message: "resize requires rows".to_owned(), + })?; + let cols = payload_u16(&frame.payload, "cols").ok_or_else(|| ErrorDto { + code: "INVALID".to_owned(), + message: "resize requires cols".to_owned(), + })?; + state + .app + .resize_terminal + .execute(ResizeTerminalInput { + session_id: sid, + rows, + cols, + }) + .map_err(ErrorDto::from) +} + +fn ws_detach( + frame: &ClientFrame, + state: &Arc, + owned: &Arc>>, +) -> Result<(), ErrorDto> { + let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?; + if let Ok(mut owned) = owned.lock() { + if let Some(index) = owned.iter().rposition(|(session, _)| *session == sid) { + let (_, gen) = owned.remove(index); + state.ws_pty_bridge.unregister_if(&sid, gen); + } + } + Ok(()) +} + +async fn ws_close( + frame: &ClientFrame, + state: &Arc, + owned: &Arc>>, +) -> Result<(SessionId, Option), ErrorDto> { + let sid = parse_session_id(&payload_string(&frame.payload, "sessionId")?)?; + ws_detach(frame, state, owned)?; + let output = state + .app + .close_terminal + .execute(CloseTerminalInput { session_id: sid }) + .await + .map_err(ErrorDto::from)?; + Ok((sid, output.code)) +} + async fn dispatch_http( method: Method, uri: Uri, @@ -825,11 +1211,447 @@ fn serialization_error(err: serde_json::Error) -> ErrorDto { } } +#[derive(Debug, Clone, Deserialize, Serialize)] +struct ClientFrame { + id: String, + kind: String, + #[serde(default)] + payload: Value, +} + +#[derive(Debug, Clone, Serialize)] +struct ServerFrame { + kind: String, + #[serde(skip_serializing_if = "Option::is_none", rename = "replyTo")] + reply_to: Option, + payload: Value, +} + +impl ServerFrame { + fn attached( + request_id: &str, + session: &TerminalSessionDto, + scrollback: Vec, + next_seq: u64, + gap: bool, + ) -> Self { + let scrollback = if scrollback.is_empty() { + Vec::new() + } else { + vec![json!({ + "seq": 0, + "bytesBase64": base64_encode(&scrollback), + })] + }; + Self { + kind: "terminal.attached".to_owned(), + reply_to: Some(request_id.to_owned()), + payload: json!({ + "session": { + "sessionId": session.session_id, + "rows": session.rows, + "cols": session.cols, + }, + "scrollback": scrollback, + "nextSeq": next_seq, + "status": "running", + "gap": gap, + "assignedConversationId": session.assigned_conversation_id, + }), + } + } + + fn output(session_id: SessionId, seq: u64, bytes: Vec) -> Self { + Self { + kind: "terminal.output".to_owned(), + reply_to: None, + payload: json!({ + "sessionId": session_id.to_string(), + "seq": seq, + "bytesBase64": base64_encode(&bytes), + }), + } + } + + fn status(session_id: SessionId, status: &str, exit_code: Option) -> Self { + Self { + kind: "terminal.status".to_owned(), + reply_to: None, + payload: json!({ + "sessionId": session_id.to_string(), + "status": status, + "exitCode": exit_code, + }), + } + } + + fn error( + request_id: Option, + session_id: Option, + code: impl Into, + message: impl Into, + ) -> Self { + Self { + kind: "error".to_owned(), + reply_to: request_id, + payload: json!({ + "sessionId": session_id, + "code": code.into(), + "message": message.into(), + }), + } + } + + fn pong() -> Self { + Self { + kind: "pong".to_owned(), + reply_to: None, + payload: json!({}), + } + } +} + +struct WsPtySink { + session_id: SessionId, + tx: mpsc::Sender, + seq: AtomicU64, +} + +impl WsPtySink { + fn new(session_id: SessionId, tx: mpsc::Sender, next_seq: u64) -> Self { + Self { + session_id, + tx, + seq: AtomicU64::new(next_seq), + } + } +} + +impl OutputSink for WsPtySink { + fn send(&self, item: PtyChunk) -> Result<(), OutputSinkError> { + let seq = self.seq.fetch_add(1, Ordering::Relaxed); + self.tx + .try_send(ServerFrame::output(self.session_id, seq, item)) + .map_err(|err| match err { + mpsc::error::TrySendError::Full(_) => OutputSinkError::Full, + mpsc::error::TrySendError::Closed(_) => OutputSinkError::Closed, + }) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum WsOpcode { + Text, + Binary, + Close, + Ping, + Pong, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct WsFrame { + opcode: WsOpcode, + payload: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum WsFrameError { + Incomplete, + Protocol(String), + TooLarge, +} + +impl std::fmt::Display for WsFrameError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Incomplete => f.write_str("incomplete websocket frame"), + Self::Protocol(message) => write!(f, "websocket protocol error: {message}"), + Self::TooLarge => f.write_str("websocket frame payload too large"), + } + } +} + +async fn read_ws_frame(reader: &mut R) -> Result +where + R: AsyncRead + Unpin, +{ + let mut header = [0_u8; 2]; + reader + .read_exact(&mut header) + .await + .map_err(|_| WsFrameError::Incomplete)?; + let len_code = header[1] & 0x7f; + let masked = header[1] & 0x80 != 0; + let mut rest = Vec::new(); + let extra_len = match len_code { + 126 => 2, + 127 => 8, + _ => 0, + }; + rest.resize(extra_len + usize::from(masked) * 4, 0); + if !rest.is_empty() { + reader + .read_exact(&mut rest) + .await + .map_err(|_| WsFrameError::Incomplete)?; + } + let mut raw = Vec::with_capacity(2 + rest.len()); + raw.extend_from_slice(&header); + raw.extend_from_slice(&rest); + let payload_len = decoded_payload_len(&raw)?; + if payload_len > WS_MAX_PAYLOAD { + return Err(WsFrameError::TooLarge); + } + let mut payload = vec![0_u8; payload_len]; + reader + .read_exact(&mut payload) + .await + .map_err(|_| WsFrameError::Incomplete)?; + raw.extend_from_slice(&payload); + decode_client_ws_frame(&raw) +} + +fn decode_client_ws_frame(raw: &[u8]) -> Result { + if raw.len() < 2 { + return Err(WsFrameError::Incomplete); + } + let fin = raw[0] & 0x80 != 0; + if !fin { + return Err(WsFrameError::Protocol( + "fragmented frames are not supported".to_owned(), + )); + } + let opcode = match raw[0] & 0x0f { + 0x1 => WsOpcode::Text, + 0x2 => WsOpcode::Binary, + 0x8 => WsOpcode::Close, + 0x9 => WsOpcode::Ping, + 0xA => WsOpcode::Pong, + other => { + return Err(WsFrameError::Protocol(format!( + "unsupported opcode {other}" + ))) + } + }; + let masked = raw[1] & 0x80 != 0; + if !masked { + return Err(WsFrameError::Protocol( + "client frames must be masked".to_owned(), + )); + } + let len_code = raw[1] & 0x7f; + let mut offset = 2; + let len = match len_code { + 126 => { + if raw.len() < offset + 2 { + return Err(WsFrameError::Incomplete); + } + let len = u16::from_be_bytes([raw[offset], raw[offset + 1]]) as usize; + offset += 2; + len + } + 127 => { + if raw.len() < offset + 8 { + return Err(WsFrameError::Incomplete); + } + let len = u64::from_be_bytes(raw[offset..offset + 8].try_into().expect("8 bytes")); + offset += 8; + usize::try_from(len).map_err(|_| WsFrameError::TooLarge)? + } + len => usize::from(len), + }; + if len > WS_MAX_PAYLOAD { + return Err(WsFrameError::TooLarge); + } + if raw.len() < offset + 4 + len { + return Err(WsFrameError::Incomplete); + } + let mask = &raw[offset..offset + 4]; + offset += 4; + let mut payload = raw[offset..offset + len].to_vec(); + for (i, byte) in payload.iter_mut().enumerate() { + *byte ^= mask[i % 4]; + } + Ok(WsFrame { opcode, payload }) +} + +fn decoded_payload_len(raw_header: &[u8]) -> Result { + let len_code = raw_header[1] & 0x7f; + let offset = 2; + let len = match len_code { + 126 => { + if raw_header.len() < offset + 2 { + return Err(WsFrameError::Incomplete); + } + u16::from_be_bytes([raw_header[offset], raw_header[offset + 1]]) as usize + } + 127 => { + if raw_header.len() < offset + 8 { + return Err(WsFrameError::Incomplete); + } + let len = + u64::from_be_bytes(raw_header[offset..offset + 8].try_into().expect("8 bytes")); + usize::try_from(len).map_err(|_| WsFrameError::TooLarge)? + } + len => usize::from(len), + }; + Ok(len) +} + +fn encode_ws_frame(opcode: WsOpcode, payload: &[u8]) -> Vec { + let opcode = match opcode { + WsOpcode::Text => 0x1, + WsOpcode::Binary => 0x2, + WsOpcode::Close => 0x8, + WsOpcode::Ping => 0x9, + WsOpcode::Pong => 0xA, + }; + let mut out = vec![0x80 | opcode]; + match payload.len() { + 0..=125 => out.push(payload.len() as u8), + 126..=65535 => { + out.push(126); + out.extend_from_slice(&(payload.len() as u16).to_be_bytes()); + } + len => { + out.push(127); + out.extend_from_slice(&(len as u64).to_be_bytes()); + } + } + out.extend_from_slice(payload); + out +} + +fn websocket_accept(key: &str) -> String { + let mut input = Vec::with_capacity(key.len() + WS_MAGIC.len()); + input.extend_from_slice(key.as_bytes()); + input.extend_from_slice(WS_MAGIC.as_bytes()); + base64_encode(&sha1_digest(&input)) +} + +fn base64_encode(bytes: &[u8]) -> String { + use base64::Engine as _; + base64::engine::general_purpose::STANDARD.encode(bytes) +} + +fn base64_decode(raw: &str) -> Result, ErrorDto> { + use base64::Engine as _; + base64::engine::general_purpose::STANDARD + .decode(raw) + .map_err(|err| ErrorDto { + code: "INVALID".to_owned(), + message: format!("invalid base64 bytes: {err}"), + }) +} + +fn sha1_digest(input: &[u8]) -> [u8; 20] { + let mut h0: u32 = 0x67452301; + let mut h1: u32 = 0xEFCDAB89; + let mut h2: u32 = 0x98BADCFE; + let mut h3: u32 = 0x10325476; + let mut h4: u32 = 0xC3D2E1F0; + + let bit_len = (input.len() as u64) * 8; + let mut msg = input.to_vec(); + msg.push(0x80); + while (msg.len() % 64) != 56 { + msg.push(0); + } + msg.extend_from_slice(&bit_len.to_be_bytes()); + + for chunk in msg.chunks(64) { + let mut w = [0_u32; 80]; + for (i, word) in w.iter_mut().take(16).enumerate() { + let j = i * 4; + *word = u32::from_be_bytes([chunk[j], chunk[j + 1], chunk[j + 2], chunk[j + 3]]); + } + for i in 16..80 { + w[i] = (w[i - 3] ^ w[i - 8] ^ w[i - 14] ^ w[i - 16]).rotate_left(1); + } + let mut a = h0; + let mut b = h1; + let mut c = h2; + let mut d = h3; + let mut e = h4; + for (i, word) in w.iter().enumerate() { + let (f, k) = match i { + 0..=19 => ((b & c) | ((!b) & d), 0x5A827999), + 20..=39 => (b ^ c ^ d, 0x6ED9EBA1), + 40..=59 => ((b & c) | (b & d) | (c & d), 0x8F1BBCDC), + _ => (b ^ c ^ d, 0xCA62C1D6), + }; + let temp = a + .rotate_left(5) + .wrapping_add(f) + .wrapping_add(e) + .wrapping_add(k) + .wrapping_add(*word); + e = d; + d = c; + c = b.rotate_left(30); + b = a; + a = temp; + } + h0 = h0.wrapping_add(a); + h1 = h1.wrapping_add(b); + h2 = h2.wrapping_add(c); + h3 = h3.wrapping_add(d); + h4 = h4.wrapping_add(e); + } + + let mut out = [0_u8; 20]; + out[0..4].copy_from_slice(&h0.to_be_bytes()); + out[4..8].copy_from_slice(&h1.to_be_bytes()); + out[8..12].copy_from_slice(&h2.to_be_bytes()); + out[12..16].copy_from_slice(&h3.to_be_bytes()); + out[16..20].copy_from_slice(&h4.to_be_bytes()); + out +} + +fn payload_string(payload: &Value, key: &str) -> Result { + payload + .get(key) + .and_then(Value::as_str) + .map(ToOwned::to_owned) + .ok_or_else(|| ErrorDto { + code: "INVALID".to_owned(), + message: format!("payload requires {key}"), + }) +} + +fn payload_u16(payload: &Value, key: &str) -> Option { + payload + .get(key) + .and_then(Value::as_u64) + .and_then(|value| u16::try_from(value).ok()) +} + +fn payload_bytes(payload: &Value) -> Result, ErrorDto> { + if let Some(raw) = payload.get("bytesBase64").and_then(Value::as_str) { + base64_decode(raw) + } else if let Some(raw) = payload.get("bytes").and_then(Value::as_str) { + base64_decode(raw) + } else { + Err(ErrorDto { + code: "INVALID".to_owned(), + message: "payload requires bytesBase64".to_owned(), + }) + } +} + +fn payload_session_id(payload: &Value) -> Option { + payload + .get("sessionId") + .and_then(Value::as_str) + .map(ToOwned::to_owned) +} + #[cfg(test)] mod tests { use super::*; use application::CreateProjectInput; use http::header::HeaderName; + use std::time::Duration; fn test_config() -> ServerConfig { ServerConfig { @@ -921,6 +1743,432 @@ mod tests { output.project.id.to_string() } + fn ws_headers(origin: &str, cookie: Option<&str>) -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert(ORIGIN, HeaderValue::from_str(origin).unwrap()); + headers.insert("upgrade", HeaderValue::from_static("websocket")); + headers.insert("connection", HeaderValue::from_static("Upgrade")); + headers.insert( + "sec-websocket-key", + HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="), + ); + if let Some(cookie) = cookie { + headers.insert(COOKIE, HeaderValue::from_str(cookie).unwrap()); + } + headers + } + + fn masked_text_frame(text: &str) -> Vec { + let payload = text.as_bytes(); + let mask = [1_u8, 2, 3, 4]; + let mut out = vec![0x81]; + out.push(0x80 | payload.len() as u8); + out.extend_from_slice(&mask); + for (i, byte) in payload.iter().enumerate() { + out.push(byte ^ mask[i % 4]); + } + out + } + + async fn recv_server_frame(rx: &mut mpsc::Receiver) -> ServerFrame { + tokio::time::timeout(Duration::from_secs(2), rx.recv()) + .await + .expect("server frame received before timeout") + .expect("server frame channel is open") + } + + fn drain_server_frames(rx: &mut mpsc::Receiver) { + while rx.try_recv().is_ok() {} + } + + fn ws_frame(id: &str, kind: &str, payload: Value) -> ClientFrame { + ClientFrame { + id: id.to_owned(), + kind: kind.to_owned(), + payload, + } + } + + fn attached_session_id(frame: &ServerFrame) -> String { + assert_eq!(frame.kind, "terminal.attached"); + frame.payload["session"]["sessionId"] + .as_str() + .expect("attached frame carries sessionId") + .to_owned() + } + + fn open_terminal_payload(command: &str, args: Vec<&str>) -> Value { + json!({ + "request": { + "cwd": std::env::temp_dir().to_string_lossy(), + "rows": 24, + "cols": 80, + "command": command, + "args": args, + } + }) + } + + async fn close_ws_session(state: &Arc, session_id: &str) { + let (tx, mut rx) = mpsc::channel(8); + let owned = Arc::new(Mutex::new(Vec::new())); + handle_client_frame( + ws_frame( + "close-cleanup", + "terminal.close", + json!({ "sessionId": session_id }), + ), + state, + &tx, + &owned, + ) + .await; + let _ = recv_server_frame(&mut rx).await; + } + + async fn wait_for_scrollback( + state: &Arc, + session_id: &str, + needle: &[u8], + ) -> Vec { + let sid = parse_session_id(session_id).expect("test session id parses"); + let handle = PtyHandle { session_id: sid }; + for _ in 0..100 { + if let Ok(scrollback) = state.app.pty_port.scrollback(&handle) { + if scrollback + .windows(needle.len()) + .any(|window| window == needle) + { + return scrollback; + } + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!( + "scrollback did not contain {:?}", + String::from_utf8_lossy(needle) + ); + } + + #[test] + fn websocket_accept_matches_rfc_example() { + assert_eq!( + websocket_accept("dGhlIHNhbXBsZSBub25jZQ=="), + "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" + ); + } + + #[test] + fn websocket_upgrade_requires_valid_cookie() { + let state = state(); + let headers = ws_headers("http://127.0.0.1:17373", None); + + let response = validate_ws_upgrade(&headers, &state).unwrap_err(); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn websocket_upgrade_rejects_invalid_cookie() { + let state = state(); + let headers = ws_headers( + "http://127.0.0.1:17373", + Some("idea_session=not-a-valid-session"), + ); + + let response = validate_ws_upgrade(&headers, &state).unwrap_err(); + + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[test] + fn websocket_upgrade_requires_allowed_origin() { + let state = state(); + let token = state.create_session(); + let headers = ws_headers( + "https://evil.example", + Some(&format!("idea_session={token}")), + ); + + let response = validate_ws_upgrade(&headers, &state).unwrap_err(); + + assert_eq!(response.status(), StatusCode::FORBIDDEN); + } + + #[test] + fn websocket_upgrade_accepts_valid_cookie_and_origin() { + let state = state(); + let token = state.create_session(); + let headers = ws_headers( + "http://127.0.0.1:17373", + Some(&format!("idea_session={token}")), + ); + + let accept = validate_ws_upgrade(&headers, &state).unwrap(); + + assert_eq!(accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + } + + #[test] + fn websocket_client_text_frame_decodes_masked_json() { + let raw = masked_text_frame(r#"{"kind":"ping"}"#); + + let frame = decode_client_ws_frame(&raw).unwrap(); + + assert_eq!(frame.opcode, WsOpcode::Text); + assert_eq!(frame.payload, br#"{"kind":"ping"}"#); + } + + #[test] + fn websocket_client_json_frame_round_trips_base64_bytes() { + let sid = domain::SessionId::new_random().to_string(); + let value = json!({ + "id": "req-1", + "kind": "input", + "payload": { + "sessionId": sid, + "bytesBase64": "AQID" + } + }); + + let frame: ClientFrame = serde_json::from_value(value.clone()).unwrap(); + let roundtrip = serde_json::to_value(&frame).unwrap(); + + assert_eq!(roundtrip, value); + assert_eq!(payload_bytes(&frame.payload).unwrap(), vec![1, 2, 3]); + } + + #[test] + fn websocket_server_attached_frame_uses_contract_shape() { + let sid = domain::SessionId::new_random().to_string(); + let frame = ServerFrame::attached( + "req-attach", + &TerminalSessionDto { + session_id: sid.clone(), + cwd: "/tmp".to_owned(), + rows: 24, + cols: 80, + assigned_conversation_id: None, + engine_session_id: None, + cell_kind: crate::dto::CellKind::Pty, + }, + vec![4, 5, 6], + 1, + true, + ); + let value = serde_json::to_value(frame).unwrap(); + + assert_eq!(value["kind"], "terminal.attached"); + assert_eq!(value["replyTo"], "req-attach"); + assert_eq!(value["payload"]["session"]["sessionId"], sid); + assert_eq!(value["payload"]["scrollback"][0]["seq"], 0); + assert_eq!(value["payload"]["scrollback"][0]["bytesBase64"], "BAUG"); + assert_eq!(value["payload"]["nextSeq"], 1); + assert_eq!(value["payload"]["gap"], true); + } + + #[tokio::test] + async fn websocket_app_ping_emits_pong() { + let state = state(); + let (tx, mut rx) = mpsc::channel(4); + let owned = Arc::new(Mutex::new(Vec::new())); + + handle_client_frame(ws_frame("ping-1", "ping", json!({})), &state, &tx, &owned).await; + let frame = recv_server_frame(&mut rx).await; + + assert_eq!(frame.kind, "pong"); + assert!(frame.payload.as_object().unwrap().is_empty()); + } + + #[tokio::test] + async fn websocket_open_terminal_emits_attached_with_empty_scrollback() { + let state = state(); + let (tx, mut rx) = mpsc::channel(16); + let owned = Arc::new(Mutex::new(Vec::new())); + + handle_client_frame( + ws_frame( + "open-1", + "terminal.open", + open_terminal_payload("/bin/sh", vec!["-c", "sleep 30"]), + ), + &state, + &tx, + &owned, + ) + .await; + let frame = recv_server_frame(&mut rx).await; + let session_id = attached_session_id(&frame); + + assert_eq!(frame.reply_to.as_deref(), Some("open-1")); + assert_eq!(frame.payload["scrollback"].as_array().unwrap().len(), 0); + assert_eq!(frame.payload["nextSeq"], 0); + + close_ws_session(&state, &session_id).await; + } + + #[tokio::test] + async fn websocket_attach_terminal_replays_scrollback_in_attached_ack() { + let state = state(); + let (tx, mut rx) = mpsc::channel(32); + let owned = Arc::new(Mutex::new(Vec::new())); + + handle_client_frame( + ws_frame( + "open-replay", + "terminal.open", + open_terminal_payload("/bin/sh", vec!["-c", "printf replay-ready; sleep 30"]), + ), + &state, + &tx, + &owned, + ) + .await; + let opened = recv_server_frame(&mut rx).await; + let session_id = attached_session_id(&opened); + + let scrollback = wait_for_scrollback(&state, &session_id, b"replay-ready").await; + assert!(scrollback.len() <= 100 * 1024); + drain_server_frames(&mut rx); + + handle_client_frame( + ws_frame( + "attach-replay", + "terminal.attach", + json!({ + "sessionId": session_id, + "lastSeq": 0, + "rows": 30, + "cols": 100, + }), + ), + &state, + &tx, + &owned, + ) + .await; + let attached = recv_server_frame(&mut rx).await; + let replay = attached.payload["scrollback"].as_array().unwrap(); + + assert_eq!(attached.kind, "terminal.attached"); + assert_eq!(attached.reply_to.as_deref(), Some("attach-replay")); + assert_eq!(attached.payload["session"]["rows"], 30); + assert_eq!(attached.payload["session"]["cols"], 100); + assert_eq!(attached.payload["gap"], true); + assert_eq!(attached.payload["nextSeq"], 1); + assert_eq!( + base64_decode(replay[0]["bytesBase64"].as_str().unwrap()).unwrap(), + scrollback + ); + + close_ws_session( + &state, + attached.payload["session"]["sessionId"].as_str().unwrap(), + ) + .await; + } + + #[tokio::test] + async fn websocket_close_terminal_emits_exited_status_and_releases_session() { + let state = state(); + let (tx, mut rx) = mpsc::channel(16); + let owned = Arc::new(Mutex::new(Vec::new())); + + handle_client_frame( + ws_frame( + "open-close", + "terminal.open", + open_terminal_payload("/bin/sh", vec!["-c", "sleep 30"]), + ), + &state, + &tx, + &owned, + ) + .await; + let opened = recv_server_frame(&mut rx).await; + let session_id = attached_session_id(&opened); + + handle_client_frame( + ws_frame( + "close-1", + "terminal.close", + json!({ "sessionId": session_id }), + ), + &state, + &tx, + &owned, + ) + .await; + let status = recv_server_frame(&mut rx).await; + + assert_eq!(status.kind, "terminal.status"); + assert_eq!(status.payload["status"], "exited"); + assert_eq!(status.payload["sessionId"], session_id); + assert_eq!(state.ws_pty_bridge.active_sessions(), 0); + assert!(state.app.terminal_sessions.is_empty()); + } + + #[test] + fn websocket_client_unmasked_frame_is_rejected() { + let raw = [0x81, 0x02, b'h', b'i']; + + let err = decode_client_ws_frame(&raw).unwrap_err(); + + assert!(matches!(err, WsFrameError::Protocol(_))); + } + + #[test] + fn websocket_payload_too_large_is_rejected() { + let len = (WS_MAX_PAYLOAD + 1) as u64; + let mut raw = vec![0x81, 0x80 | 127]; + raw.extend_from_slice(&len.to_be_bytes()); + raw.extend_from_slice(&[1, 2, 3, 4]); + + let err = decode_client_ws_frame(&raw).unwrap_err(); + + assert_eq!(err, WsFrameError::TooLarge); + } + + #[test] + fn websocket_server_frames_are_unmasked() { + let raw = encode_ws_frame(WsOpcode::Text, b"{}"); + + assert_eq!(raw[0], 0x81); + assert_eq!(raw[1] & 0x80, 0, "server frames must not be masked"); + } + + #[tokio::test] + async fn websocket_sink_full_reports_backpressure_without_blocking() { + let (tx, mut rx) = mpsc::channel(1); + let sid = domain::SessionId::new_random(); + let sink = WsPtySink::new(sid, tx, 0); + + assert!(sink.send(vec![1]).is_ok()); + assert_eq!(sink.send(vec![2]), Err(OutputSinkError::Full)); + let frame = rx.recv().await.unwrap(); + assert_eq!(frame.kind, "terminal.output"); + } + + #[tokio::test] + async fn websocket_output_bridge_replaces_attachment_without_double_delivery() { + let bridge = OutputBridge::::new(); + let sid = domain::SessionId::new_random(); + let (old_tx, mut old_rx) = mpsc::channel(4); + let (new_tx, mut new_rx) = mpsc::channel(4); + + let old_gen = bridge.register(sid, Arc::new(WsPtySink::new(sid, old_tx, 0))); + let _new_gen = bridge.register(sid, Arc::new(WsPtySink::new(sid, new_tx, 0))); + + assert!(bridge.send_output(&sid, vec![9])); + assert!(old_rx.try_recv().is_err()); + let frame = new_rx.recv().await.unwrap(); + assert_eq!(frame.kind, "terminal.output"); + assert_eq!(frame.payload["bytesBase64"], "CQ=="); + + bridge.unregister_if(&sid, old_gen); + assert_eq!(bridge.active_sessions(), 1); + } + #[test] fn public_bind_requires_explicit_remote_security() { let config = ServerConfig {