diff --git a/Cargo.lock b/Cargo.lock index 1377c479..85dfeab3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -343,6 +343,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b52af3cb4058c895d37317bb27508dccc8e5f2d39454016b297bf4a400597b8" dependencies = [ "axum-core", + "base64 0.22.1", "bytes", "form_urlencoded", "futures-util", @@ -361,8 +362,10 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_urlencoded", + "sha1", "sync_wrapper", "tokio", + "tokio-tungstenite 0.28.0", "tower", "tower-layer", "tower-service", @@ -1147,6 +1150,12 @@ dependencies = [ "syn 2.0.114", ] +[[package]] +name = "data-encoding" +version = "2.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7a1e2f27636f116493b8b860f5546edb47c8d8f8ea73e1d2a20be88e28d1fea" + [[package]] name = "deadpool" version = "0.12.3" @@ -2214,6 +2223,7 @@ dependencies = [ "tokio-postgres", "tokio-stream", "tokio-test", + "tokio-tungstenite 0.26.2", "tower", "tower-http", "tracing", @@ -4376,6 +4386,30 @@ dependencies = [ "tokio-stream", ] +[[package]] +name = "tokio-tungstenite" +version = "0.26.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7a9daff607c6d2bf6c16fd681ccb7eecc83e4e2cdc1ca067ffaadfca5de7f084" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite 0.26.2", +] + +[[package]] +name = "tokio-tungstenite" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d25a406cddcc431a75d3d9afc6a7c0f7428d4891dd973e4d54c56b46127bf857" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite 0.28.0", +] + [[package]] name = "tokio-util" version = "0.7.18" @@ -4588,6 +4622,40 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "tungstenite" +version = "0.26.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4793cb5e56680ecbb1d843515b23b6de9a75eb04b66643e256a396d43be33c13" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.2", + "sha1", + "thiserror 2.0.18", + "utf-8", +] + +[[package]] +name = "tungstenite" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8628dcc84e5a09eb3d8423d6cb682965dea9133204e8fb3efee74c2a0c259442" +dependencies = [ + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.9.2", + "sha1", + "thiserror 2.0.18", + "utf-8", +] + [[package]] name = "typenum" version = "1.19.0" @@ -4691,6 +4759,12 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "daf8dba3b7eb870caf1ddeed7bc9d2a049f3cfdfae7cb521b087cc33ae4c49da" +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/Cargo.toml b/Cargo.toml index 94177005..6fc39003 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -54,7 +54,7 @@ rustyline = { version = "17", features = ["derive", "with-file-history"] } termimad = "0.34" # Channel integrations -axum = "0.8" +axum = { version = "0.8", features = ["ws"] } tower = "0.5" tower-http = { version = "0.6", features = ["trace", "cors"] } @@ -111,6 +111,7 @@ zbus = "4" [dev-dependencies] tokio-test = "0.4" +tokio-tungstenite = "0.26" testcontainers-modules = { version = "0.11", features = ["postgres"] } pretty_assertions = "1" tempfile = "3" diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 64fc9611..2eb23c86 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -8,6 +8,7 @@ //! ```text //! Browser ─── POST /api/chat/send ──► Agent Loop //! ◄── GET /api/chat/events ── SSE stream +//! ─── GET /api/chat/ws ─────► WebSocket (bidirectional) //! ─── GET /api/memory/* ────► Workspace //! ─── GET /api/jobs/* ──────► ContextManager //! ◄── GET / ───────────────── Static HTML/CSS/JS @@ -18,6 +19,7 @@ pub mod log_layer; pub mod server; pub mod sse; pub mod types; +pub mod ws; use std::net::SocketAddr; use std::sync::Arc; @@ -75,6 +77,7 @@ impl GatewayChannel { tool_registry: None, user_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), }); Self { @@ -97,6 +100,7 @@ impl GatewayChannel { tool_registry: self.state.tool_registry.clone(), user_id: self.state.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: self.state.ws_tracker.clone(), }; mutate(&mut new_state); self.state = Arc::new(new_state); @@ -169,9 +173,10 @@ impl Channel for GatewayChannel { ), })?; - server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?; + let bound_addr = + server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?; - tracing::info!("Web gateway listening on http://{}", addr); + tracing::info!("Web gateway listening on http://{}", bound_addr); tracing::info!("Auth token: {}", self.auth_token); Ok(Box::pin(ReceiverStream::new(rx))) diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index c862ba4c..11363dfa 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -8,7 +8,7 @@ use std::sync::Arc; use axum::{ Json, Router, - extract::{Path, Query, State}, + extract::{Path, Query, State, WebSocketUpgrade}, http::{StatusCode, header}, middleware, response::{ @@ -55,20 +55,31 @@ pub struct GatewayState { pub user_id: String, /// Shutdown signal sender. pub shutdown_tx: tokio::sync::RwLock>>, + /// WebSocket connection tracker. + pub ws_tracker: Option>, } /// Start the gateway HTTP server. +/// +/// Returns the actual bound `SocketAddr` (useful when binding to port 0). pub async fn start_server( addr: SocketAddr, state: Arc, auth_token: String, -) -> Result<(), crate::error::ChannelError> { +) -> Result { let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { crate::error::ChannelError::StartupFailed { name: "gateway".to_string(), reason: format!("Failed to bind to {}: {}", addr, e), } })?; + let bound_addr = + listener + .local_addr() + .map_err(|e| crate::error::ChannelError::StartupFailed { + name: "gateway".to_string(), + reason: format!("Failed to get local addr: {}", e), + })?; // Public routes (no auth) let public = Router::new().route("/api/health", get(health_handler)); @@ -80,6 +91,7 @@ pub async fn start_server( .route("/api/chat/send", post(chat_send_handler)) .route("/api/chat/approval", post(chat_approval_handler)) .route("/api/chat/events", get(chat_events_handler)) + .route("/api/chat/ws", get(chat_ws_handler)) .route("/api/chat/history", get(chat_history_handler)) .route("/api/chat/threads", get(chat_threads_handler)) .route("/api/chat/thread/new", post(chat_new_thread_handler)) @@ -108,6 +120,8 @@ pub async fn start_server( "/api/extensions/{name}/remove", post(extensions_remove_handler), ) + // Gateway control plane + .route("/api/gateway/status", get(gateway_status_handler)) .route_layer(middleware::from_fn_with_state(auth_state, auth_middleware)); // Static file routes (no auth, served from embedded strings) @@ -137,7 +151,7 @@ pub async fn start_server( } }); - Ok(()) + Ok(bound_addr) } // --- Static file handlers --- @@ -272,6 +286,13 @@ async fn chat_events_handler(State(state): State>) -> impl Int state.sse.subscribe() } +async fn chat_ws_handler( + ws: WebSocketUpgrade, + State(state): State>, +) -> impl IntoResponse { + ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state)) +} + #[derive(Deserialize)] struct HistoryQuery { thread_id: Option, @@ -834,3 +855,29 @@ async fn extensions_remove_handler( Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } } + +// --- Gateway control plane handlers --- + +async fn gateway_status_handler( + State(state): State>, +) -> Json { + let sse_connections = state.sse.connection_count(); + let ws_connections = state + .ws_tracker + .as_ref() + .map(|t| t.connection_count()) + .unwrap_or(0); + + Json(GatewayStatusResponse { + sse_connections, + ws_connections, + total_connections: sse_connections + ws_connections, + }) +} + +#[derive(serde::Serialize)] +struct GatewayStatusResponse { + sse_connections: u64, + ws_connections: u64, + total_connections: u64, +} diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index e15446e1..132748fe 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -41,6 +41,23 @@ impl SseManager { self.connection_count.load(Ordering::Relaxed) } + /// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket). + /// + /// Returns a stream of `SseEvent` values and increments/decrements the + /// connection counter on creation/drop, just like `subscribe()` does for SSE. + pub fn subscribe_raw(&self) -> impl Stream + Send + 'static + use<> { + let counter = Arc::clone(&self.connection_count); + counter.fetch_add(1, Ordering::Relaxed); + let rx = self.tx.subscribe(); + + let stream = BroadcastStream::new(rx).filter_map(|result| result.ok()); + + CountedStream { + inner: stream, + counter, + } + } + /// Create a new SSE stream for a client connection. pub fn subscribe( &self, @@ -144,4 +161,53 @@ mod tests { _ => panic!("unexpected event type"), } } + + #[tokio::test] + async fn test_subscribe_raw_receives_events() { + let manager = SseManager::new(); + let mut stream = Box::pin(manager.subscribe_raw()); + + assert_eq!(manager.connection_count(), 1); + + manager.broadcast(SseEvent::Thinking { + message: "working".to_string(), + }); + + let event = stream.next().await.unwrap(); + match event { + SseEvent::Thinking { message } => assert_eq!(message, "working"), + _ => panic!("Expected Thinking event"), + } + } + + #[tokio::test] + async fn test_subscribe_raw_decrements_on_drop() { + let manager = SseManager::new(); + { + let _stream = Box::pin(manager.subscribe_raw()); + assert_eq!(manager.connection_count(), 1); + } + // Stream dropped, counter should decrement + assert_eq!(manager.connection_count(), 0); + } + + #[tokio::test] + async fn test_subscribe_raw_multiple_subscribers() { + let manager = SseManager::new(); + let mut s1 = Box::pin(manager.subscribe_raw()); + let mut s2 = Box::pin(manager.subscribe_raw()); + assert_eq!(manager.connection_count(), 2); + + manager.broadcast(SseEvent::Heartbeat); + + let e1 = s1.next().await.unwrap(); + let e2 = s2.next().await.unwrap(); + assert!(matches!(e1, SseEvent::Heartbeat)); + assert!(matches!(e2, SseEvent::Heartbeat)); + + drop(s1); + assert_eq!(manager.connection_count(), 1); + drop(s2); + assert_eq!(manager.connection_count(), 0); + } } diff --git a/src/channels/web/types.rs b/src/channels/web/types.rs index 1159acb4..ee5e6122 100644 --- a/src/channels/web/types.rs +++ b/src/channels/web/types.rs @@ -260,6 +260,72 @@ impl ActionResponse { } } +// --- WebSocket --- + +/// Message sent by a WebSocket client to the server. +#[derive(Debug, Clone, Deserialize)] +#[serde(tag = "type")] +pub enum WsClientMessage { + /// Send a chat message to the agent. + #[serde(rename = "message")] + Message { + content: String, + thread_id: Option, + }, + /// Approve or deny a pending tool execution. + #[serde(rename = "approval")] + Approval { + request_id: String, + /// "approve", "always", or "deny" + action: String, + }, + /// Client heartbeat ping. + #[serde(rename = "ping")] + Ping, +} + +/// Message sent by the server to a WebSocket client. +#[derive(Debug, Clone, Serialize)] +#[serde(tag = "type")] +pub enum WsServerMessage { + /// An SSE-style event forwarded over WebSocket. + #[serde(rename = "event")] + Event { + /// The event sub-type (response, thinking, tool_started, etc.) + event_type: String, + /// The event payload as a JSON value. + data: serde_json::Value, + }, + /// Server heartbeat pong. + #[serde(rename = "pong")] + Pong, + /// Error message. + #[serde(rename = "error")] + Error { message: String }, +} + +impl WsServerMessage { + /// Create a WsServerMessage from an SseEvent. + pub fn from_sse_event(event: &SseEvent) -> Self { + let event_type = match event { + SseEvent::Response { .. } => "response", + SseEvent::Thinking { .. } => "thinking", + SseEvent::ToolStarted { .. } => "tool_started", + SseEvent::ToolCompleted { .. } => "tool_completed", + SseEvent::StreamChunk { .. } => "stream_chunk", + SseEvent::Status { .. } => "status", + SseEvent::ApprovalNeeded { .. } => "approval_needed", + SseEvent::Error { .. } => "error", + SseEvent::Heartbeat => "heartbeat", + }; + let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null); + WsServerMessage::Event { + event_type: event_type.to_string(), + data, + } + } +} + // --- Health --- #[derive(Debug, Serialize)] @@ -267,3 +333,145 @@ pub struct HealthResponse { pub status: &'static str, pub channel: &'static str, } + +#[cfg(test)] +mod tests { + use super::*; + + // ---- WsClientMessage deserialization tests ---- + + #[test] + fn test_ws_client_message_parse() { + let json = r#"{"type":"message","content":"hello","thread_id":"t1"}"#; + let msg: WsClientMessage = serde_json::from_str(json).unwrap(); + match msg { + WsClientMessage::Message { content, thread_id } => { + assert_eq!(content, "hello"); + assert_eq!(thread_id.as_deref(), Some("t1")); + } + _ => panic!("Expected Message variant"), + } + } + + #[test] + fn test_ws_client_message_no_thread() { + let json = r#"{"type":"message","content":"hi"}"#; + let msg: WsClientMessage = serde_json::from_str(json).unwrap(); + match msg { + WsClientMessage::Message { content, thread_id } => { + assert_eq!(content, "hi"); + assert!(thread_id.is_none()); + } + _ => panic!("Expected Message variant"), + } + } + + #[test] + fn test_ws_client_approval_parse() { + let json = r#"{"type":"approval","request_id":"abc-123","action":"approve"}"#; + let msg: WsClientMessage = serde_json::from_str(json).unwrap(); + match msg { + WsClientMessage::Approval { request_id, action } => { + assert_eq!(request_id, "abc-123"); + assert_eq!(action, "approve"); + } + _ => panic!("Expected Approval variant"), + } + } + + #[test] + fn test_ws_client_ping_parse() { + let json = r#"{"type":"ping"}"#; + let msg: WsClientMessage = serde_json::from_str(json).unwrap(); + assert!(matches!(msg, WsClientMessage::Ping)); + } + + #[test] + fn test_ws_client_unknown_type_fails() { + let json = r#"{"type":"unknown"}"#; + let result: Result = serde_json::from_str(json); + assert!(result.is_err()); + } + + // ---- WsServerMessage serialization tests ---- + + #[test] + fn test_ws_server_pong_serialize() { + let msg = WsServerMessage::Pong; + let json = serde_json::to_string(&msg).unwrap(); + assert_eq!(json, r#"{"type":"pong"}"#); + } + + #[test] + fn test_ws_server_error_serialize() { + let msg = WsServerMessage::Error { + message: "bad request".to_string(), + }; + let json = serde_json::to_string(&msg).unwrap(); + let parsed: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed["type"], "error"); + assert_eq!(parsed["message"], "bad request"); + } + + #[test] + fn test_ws_server_from_sse_response() { + let sse = SseEvent::Response { + content: "hello".to_string(), + thread_id: "t1".to_string(), + }; + let ws = WsServerMessage::from_sse_event(&sse); + match ws { + WsServerMessage::Event { event_type, data } => { + assert_eq!(event_type, "response"); + assert_eq!(data["content"], "hello"); + assert_eq!(data["thread_id"], "t1"); + } + _ => panic!("Expected Event variant"), + } + } + + #[test] + fn test_ws_server_from_sse_thinking() { + let sse = SseEvent::Thinking { + message: "reasoning...".to_string(), + }; + let ws = WsServerMessage::from_sse_event(&sse); + match ws { + WsServerMessage::Event { event_type, data } => { + assert_eq!(event_type, "thinking"); + assert_eq!(data["message"], "reasoning..."); + } + _ => panic!("Expected Event variant"), + } + } + + #[test] + fn test_ws_server_from_sse_approval_needed() { + let sse = SseEvent::ApprovalNeeded { + request_id: "r1".to_string(), + tool_name: "shell".to_string(), + description: "Run ls".to_string(), + parameters: "{}".to_string(), + }; + let ws = WsServerMessage::from_sse_event(&sse); + match ws { + WsServerMessage::Event { event_type, data } => { + assert_eq!(event_type, "approval_needed"); + assert_eq!(data["tool_name"], "shell"); + } + _ => panic!("Expected Event variant"), + } + } + + #[test] + fn test_ws_server_from_sse_heartbeat() { + let sse = SseEvent::Heartbeat; + let ws = WsServerMessage::from_sse_event(&sse); + match ws { + WsServerMessage::Event { event_type, .. } => { + assert_eq!(event_type, "heartbeat"); + } + _ => panic!("Expected Event variant"), + } + } +} diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs new file mode 100644 index 00000000..64755e3c --- /dev/null +++ b/src/channels/web/ws.rs @@ -0,0 +1,411 @@ +//! WebSocket handler for bidirectional client communication. +//! +//! Provides the same event stream as SSE but also accepts incoming messages +//! (chat, approvals) over a single persistent connection. +//! +//! ```text +//! Client ──── WS frame: {"type":"message","content":"hello"} ──► Agent Loop +//! ◄─── WS frame: {"type":"event","event_type":"response","data":{...}} ── Broadcast +//! ──── WS frame: {"type":"ping"} ──────────────────────────────────────► +//! ◄─── WS frame: {"type":"pong"} ────────────────────────────────────── +//! ``` + +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; + +use axum::extract::ws::{Message, WebSocket}; +use futures::{SinkExt, StreamExt}; +use tokio::sync::mpsc; +use uuid::Uuid; + +use crate::agent::submission::Submission; +use crate::channels::IncomingMessage; +use crate::channels::web::server::GatewayState; +use crate::channels::web::types::{WsClientMessage, WsServerMessage}; + +/// Tracks active WebSocket connections. +pub struct WsConnectionTracker { + count: AtomicU64, +} + +impl WsConnectionTracker { + pub fn new() -> Self { + Self { + count: AtomicU64::new(0), + } + } + + pub fn connection_count(&self) -> u64 { + self.count.load(Ordering::Relaxed) + } + + fn increment(&self) { + self.count.fetch_add(1, Ordering::Relaxed); + } + + fn decrement(&self) { + self.count.fetch_sub(1, Ordering::Relaxed); + } +} + +impl Default for WsConnectionTracker { + fn default() -> Self { + Self::new() + } +} + +/// Handle an upgraded WebSocket connection. +/// +/// Spawns two tasks: +/// - **sender**: forwards broadcast events to the WebSocket client +/// - **receiver**: reads client frames and routes them to the agent +/// +/// When either task ends (client disconnect or broadcast closed), both are +/// cleaned up. +pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { + let (mut ws_sink, mut ws_stream) = socket.split(); + + // Track connection + if let Some(ref tracker) = state.ws_tracker { + tracker.increment(); + } + let tracker_for_drop = state.ws_tracker.clone(); + + // Subscribe to broadcast events (same source as SSE) + let mut event_stream = Box::pin(state.sse.subscribe_raw()); + + // Channel for the sender task to receive messages from both + // the broadcast stream and any direct sends (like Pong) + let (direct_tx, mut direct_rx) = mpsc::channel::(64); + + // Sender task: forward broadcast events + direct messages to WS client + let sender_handle = tokio::spawn(async move { + loop { + let msg = tokio::select! { + event = event_stream.next() => { + match event { + Some(sse_event) => WsServerMessage::from_sse_event(&sse_event), + None => break, // Broadcast channel closed + } + } + direct = direct_rx.recv() => { + match direct { + Some(msg) => msg, + None => break, // Direct channel closed + } + } + }; + + let json = match serde_json::to_string(&msg) { + Ok(j) => j, + Err(_) => continue, + }; + + if ws_sink.send(Message::Text(json.into())).await.is_err() { + break; // Client disconnected + } + } + }); + + // Receiver task: read client frames and route to agent + let user_id = state.user_id.clone(); + while let Some(Ok(frame)) = ws_stream.next().await { + match frame { + Message::Text(text) => { + let parsed: Result = serde_json::from_str(&text); + match parsed { + Ok(client_msg) => { + handle_client_message(client_msg, &state, &user_id, &direct_tx).await; + } + Err(e) => { + let _ = direct_tx + .send(WsServerMessage::Error { + message: format!("Invalid message: {}", e), + }) + .await; + } + } + } + Message::Close(_) => break, + // Ignore binary, ping/pong (axum handles protocol-level pings) + _ => {} + } + } + + // Clean up: abort sender, decrement counter + sender_handle.abort(); + if let Some(ref tracker) = tracker_for_drop { + tracker.decrement(); + } +} + +/// Route a parsed client message to the appropriate handler. +async fn handle_client_message( + msg: WsClientMessage, + state: &GatewayState, + user_id: &str, + direct_tx: &mpsc::Sender, +) { + match msg { + WsClientMessage::Message { content, thread_id } => { + let mut incoming = IncomingMessage::new("gateway", user_id, &content); + if let Some(ref tid) = thread_id { + incoming = incoming.with_thread(tid); + } + + let tx_guard = state.msg_tx.read().await; + if let Some(ref tx) = *tx_guard { + if tx.send(incoming).await.is_err() { + let _ = direct_tx + .send(WsServerMessage::Error { + message: "Channel closed".to_string(), + }) + .await; + } + } else { + let _ = direct_tx + .send(WsServerMessage::Error { + message: "Channel not started".to_string(), + }) + .await; + } + } + WsClientMessage::Approval { request_id, action } => { + let (approved, always) = match action.as_str() { + "approve" => (true, false), + "always" => (true, true), + "deny" => (false, false), + other => { + let _ = direct_tx + .send(WsServerMessage::Error { + message: format!("Unknown approval action: {}", other), + }) + .await; + return; + } + }; + + let request_uuid = match Uuid::parse_str(&request_id) { + Ok(id) => id, + Err(_) => { + let _ = direct_tx + .send(WsServerMessage::Error { + message: "Invalid request_id (expected UUID)".to_string(), + }) + .await; + return; + } + }; + + let approval = Submission::ExecApproval { + request_id: request_uuid, + approved, + always, + }; + let content = match serde_json::to_string(&approval) { + Ok(c) => c, + Err(e) => { + let _ = direct_tx + .send(WsServerMessage::Error { + message: format!("Failed to serialize approval: {}", e), + }) + .await; + return; + } + }; + + let msg = IncomingMessage::new("gateway", user_id, content); + let tx_guard = state.msg_tx.read().await; + if let Some(ref tx) = *tx_guard { + let _ = tx.send(msg).await; + } + } + WsClientMessage::Ping => { + let _ = direct_tx.send(WsServerMessage::Pong).await; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_ws_connection_tracker() { + let tracker = WsConnectionTracker::new(); + assert_eq!(tracker.connection_count(), 0); + + tracker.increment(); + assert_eq!(tracker.connection_count(), 1); + + tracker.increment(); + assert_eq!(tracker.connection_count(), 2); + + tracker.decrement(); + assert_eq!(tracker.connection_count(), 1); + + tracker.decrement(); + assert_eq!(tracker.connection_count(), 0); + } + + #[test] + fn test_ws_connection_tracker_default() { + let tracker = WsConnectionTracker::default(); + assert_eq!(tracker.connection_count(), 0); + } + + #[tokio::test] + async fn test_handle_client_message_ping() { + // Ping should produce a Pong on the direct channel + let (direct_tx, mut direct_rx) = mpsc::channel(16); + let state = make_test_state(None).await; + + handle_client_message(WsClientMessage::Ping, &state, "user1", &direct_tx).await; + + let response = direct_rx.recv().await.unwrap(); + assert!(matches!(response, WsServerMessage::Pong)); + } + + #[tokio::test] + async fn test_handle_client_message_sends_to_agent() { + // A Message should be forwarded to the agent's msg_tx + let (agent_tx, mut agent_rx) = mpsc::channel(16); + let state = make_test_state(Some(agent_tx)).await; + let (direct_tx, _direct_rx) = mpsc::channel(16); + + handle_client_message( + WsClientMessage::Message { + content: "hello agent".to_string(), + thread_id: Some("t1".to_string()), + }, + &state, + "user1", + &direct_tx, + ) + .await; + + let incoming = agent_rx.recv().await.unwrap(); + assert_eq!(incoming.content, "hello agent"); + assert_eq!(incoming.thread_id.as_deref(), Some("t1")); + assert_eq!(incoming.channel, "gateway"); + assert_eq!(incoming.user_id, "user1"); + } + + #[tokio::test] + async fn test_handle_client_message_no_channel() { + // When msg_tx is None, should send an error back + let state = make_test_state(None).await; + let (direct_tx, mut direct_rx) = mpsc::channel(16); + + handle_client_message( + WsClientMessage::Message { + content: "hello".to_string(), + thread_id: None, + }, + &state, + "user1", + &direct_tx, + ) + .await; + + let response = direct_rx.recv().await.unwrap(); + match response { + WsServerMessage::Error { message } => { + assert!(message.contains("not started")); + } + _ => panic!("Expected Error variant"), + } + } + + #[tokio::test] + async fn test_handle_client_approval_approve() { + let (agent_tx, mut agent_rx) = mpsc::channel(16); + let state = make_test_state(Some(agent_tx)).await; + let (direct_tx, _direct_rx) = mpsc::channel(16); + + let request_id = Uuid::new_v4(); + handle_client_message( + WsClientMessage::Approval { + request_id: request_id.to_string(), + action: "approve".to_string(), + }, + &state, + "user1", + &direct_tx, + ) + .await; + + let incoming = agent_rx.recv().await.unwrap(); + // The content should be a serialized ExecApproval + assert!(incoming.content.contains("ExecApproval")); + } + + #[tokio::test] + async fn test_handle_client_approval_invalid_action() { + let state = make_test_state(None).await; + let (direct_tx, mut direct_rx) = mpsc::channel(16); + + handle_client_message( + WsClientMessage::Approval { + request_id: Uuid::new_v4().to_string(), + action: "maybe".to_string(), + }, + &state, + "user1", + &direct_tx, + ) + .await; + + let response = direct_rx.recv().await.unwrap(); + match response { + WsServerMessage::Error { message } => { + assert!(message.contains("Unknown approval action")); + } + _ => panic!("Expected Error variant"), + } + } + + #[tokio::test] + async fn test_handle_client_approval_invalid_uuid() { + let state = make_test_state(None).await; + let (direct_tx, mut direct_rx) = mpsc::channel(16); + + handle_client_message( + WsClientMessage::Approval { + request_id: "not-a-uuid".to_string(), + action: "approve".to_string(), + }, + &state, + "user1", + &direct_tx, + ) + .await; + + let response = direct_rx.recv().await.unwrap(); + match response { + WsServerMessage::Error { message } => { + assert!(message.contains("Invalid request_id")); + } + _ => panic!("Expected Error variant"), + } + } + + /// Helper to create a GatewayState for testing. + async fn make_test_state(msg_tx: Option>) -> GatewayState { + use crate::channels::web::sse::SseManager; + + GatewayState { + msg_tx: tokio::sync::RwLock::new(msg_tx), + sse: SseManager::new(), + workspace: None, + context_manager: None, + session_manager: None, + log_broadcaster: None, + extension_manager: None, + tool_registry: None, + user_id: "test".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + } + } +} diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs new file mode 100644 index 00000000..a93bc49d --- /dev/null +++ b/tests/ws_gateway_integration.rs @@ -0,0 +1,319 @@ +//! End-to-end integration tests for the WebSocket gateway. +//! +//! These tests start a real Axum server on a random port, connect a WebSocket +//! client, and verify the full message flow: +//! - WebSocket upgrade with auth +//! - Ping/pong +//! - Client message → agent msg_tx +//! - Broadcast SSE event → WebSocket client +//! - Connection tracking (counter increment/decrement) +//! - Gateway status endpoint + +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use futures::{SinkExt, StreamExt}; +use tokio::sync::mpsc; +use tokio::time::timeout; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; + +use ironclaw::channels::IncomingMessage; +use ironclaw::channels::web::server::{GatewayState, start_server}; +use ironclaw::channels::web::sse::SseManager; +use ironclaw::channels::web::types::SseEvent; +use ironclaw::channels::web::ws::WsConnectionTracker; + +const AUTH_TOKEN: &str = "test-token-12345"; +const TIMEOUT: Duration = Duration::from_secs(5); + +/// Start a gateway server on a random port and return the bound address + agent +/// message receiver. +async fn start_test_server() -> ( + SocketAddr, + Arc, + mpsc::Receiver, +) { + let (agent_tx, agent_rx) = mpsc::channel(64); + + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), + sse: SseManager::new(), + workspace: None, + context_manager: None, + session_manager: None, + log_broadcaster: None, + extension_manager: None, + tool_registry: None, + user_id: "test-user".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + }); + + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + .await + .expect("Failed to start test server"); + + (bound_addr, state, agent_rx) +} + +/// Connect a WebSocket client with auth token in query parameter. +async fn connect_ws( + addr: SocketAddr, +) -> tokio_tungstenite::WebSocketStream> { + let url = format!("ws://{}/api/chat/ws?token={}", addr, AUTH_TOKEN); + let request = url.into_client_request().unwrap(); + let (stream, _response) = tokio_tungstenite::connect_async(request) + .await + .expect("Failed to connect WebSocket"); + stream +} + +/// Read the next text frame from the WebSocket, with a timeout. +async fn recv_text( + stream: &mut (impl StreamExt> + Unpin), +) -> String { + let msg = timeout(TIMEOUT, stream.next()) + .await + .expect("Timed out waiting for WS message") + .expect("Stream ended") + .expect("WS error"); + match msg { + Message::Text(text) => text.to_string(), + other => panic!("Expected Text frame, got {:?}", other), + } +} + +// ============================================================================ +// Tests +// ============================================================================ + +#[tokio::test] +async fn test_ws_ping_pong() { + let (addr, _state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + + // Send ping + let ping = r#"{"type":"ping"}"#; + ws.send(Message::Text(ping.into())).await.unwrap(); + + // Expect pong + let text = recv_text(&mut ws).await; + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "pong"); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_ws_message_reaches_agent() { + let (addr, _state, mut agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + + // Send a chat message + let msg = r#"{"type":"message","content":"hello from ws","thread_id":"t42"}"#; + ws.send(Message::Text(msg.into())).await.unwrap(); + + // Verify it arrives on the agent's msg_tx + let incoming = timeout(TIMEOUT, agent_rx.recv()) + .await + .expect("Timed out waiting for agent message") + .expect("Agent channel closed"); + + assert_eq!(incoming.content, "hello from ws"); + assert_eq!(incoming.thread_id.as_deref(), Some("t42")); + assert_eq!(incoming.channel, "gateway"); + assert_eq!(incoming.user_id, "test-user"); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_ws_broadcast_event_received() { + let (addr, state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + + // Give the connection a moment to fully establish + tokio::time::sleep(Duration::from_millis(50)).await; + + // Broadcast an SSE event (simulates agent sending a response) + state.sse.broadcast(SseEvent::Response { + content: "agent says hi".to_string(), + thread_id: "t1".to_string(), + }); + + // The WS client should receive it + let text = recv_text(&mut ws).await; + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!(parsed["event_type"], "response"); + assert_eq!(parsed["data"]["content"], "agent says hi"); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_ws_thinking_event() { + let (addr, state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + tokio::time::sleep(Duration::from_millis(50)).await; + + state.sse.broadcast(SseEvent::Thinking { + message: "analyzing...".to_string(), + }); + + let text = recv_text(&mut ws).await; + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!(parsed["event_type"], "thinking"); + assert_eq!(parsed["data"]["message"], "analyzing..."); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_ws_connection_tracking() { + let (addr, state, _agent_rx) = start_test_server().await; + let tracker = state.ws_tracker.as_ref().unwrap(); + + assert_eq!(tracker.connection_count(), 0); + + // Connect first client + let ws1 = connect_ws(addr).await; + tokio::time::sleep(Duration::from_millis(50)).await; + assert_eq!(tracker.connection_count(), 1); + + // Connect second client + let ws2 = connect_ws(addr).await; + tokio::time::sleep(Duration::from_millis(50)).await; + assert_eq!(tracker.connection_count(), 2); + + // Disconnect first + drop(ws1); + tokio::time::sleep(Duration::from_millis(100)).await; + assert_eq!(tracker.connection_count(), 1); + + // Disconnect second + drop(ws2); + tokio::time::sleep(Duration::from_millis(100)).await; + assert_eq!(tracker.connection_count(), 0); +} + +#[tokio::test] +async fn test_ws_invalid_message_returns_error() { + let (addr, _state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + + // Send invalid JSON + ws.send(Message::Text("not json".into())).await.unwrap(); + + // Should get an error message back + let text = recv_text(&mut ws).await; + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "error"); + assert!( + parsed["message"] + .as_str() + .unwrap() + .contains("Invalid message") + ); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_ws_unknown_type_returns_error() { + let (addr, _state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + + // Send valid JSON but unknown message type + ws.send(Message::Text(r#"{"type":"foobar"}"#.into())) + .await + .unwrap(); + + let text = recv_text(&mut ws).await; + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "error"); + + ws.close(None).await.unwrap(); +} + +#[tokio::test] +async fn test_gateway_status_endpoint() { + let (addr, _state, _agent_rx) = start_test_server().await; + + // Connect a WS client + let _ws = connect_ws(addr).await; + tokio::time::sleep(Duration::from_millis(50)).await; + + // Hit the status endpoint + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", format!("Bearer {}", AUTH_TOKEN)) + .send() + .await + .expect("Failed to fetch status"); + + assert_eq!(resp.status(), 200); + + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["ws_connections"], 1); + assert!(body["total_connections"].as_u64().unwrap() >= 1); +} + +#[tokio::test] +async fn test_ws_no_auth_rejected() { + let (addr, _state, _agent_rx) = start_test_server().await; + + // Try to connect without auth token + let url = format!("ws://{}/api/chat/ws", addr); + let request = url.into_client_request().unwrap(); + let result = tokio_tungstenite::connect_async(request).await; + + // Should fail (401 from auth middleware before WS upgrade) + assert!(result.is_err()); +} + +#[tokio::test] +async fn test_ws_multiple_events_in_sequence() { + let (addr, state, _agent_rx) = start_test_server().await; + let mut ws = connect_ws(addr).await; + tokio::time::sleep(Duration::from_millis(50)).await; + + // Broadcast multiple events rapidly + state.sse.broadcast(SseEvent::Thinking { + message: "step 1".to_string(), + }); + state.sse.broadcast(SseEvent::ToolStarted { + name: "shell".to_string(), + }); + state.sse.broadcast(SseEvent::ToolCompleted { + name: "shell".to_string(), + success: true, + }); + state.sse.broadcast(SseEvent::Response { + content: "done".to_string(), + thread_id: "t1".to_string(), + }); + + // Receive all 4 in order + let t1 = recv_text(&mut ws).await; + let t2 = recv_text(&mut ws).await; + let t3 = recv_text(&mut ws).await; + let t4 = recv_text(&mut ws).await; + + let p1: serde_json::Value = serde_json::from_str(&t1).unwrap(); + let p2: serde_json::Value = serde_json::from_str(&t2).unwrap(); + let p3: serde_json::Value = serde_json::from_str(&t3).unwrap(); + let p4: serde_json::Value = serde_json::from_str(&t4).unwrap(); + + assert_eq!(p1["event_type"], "thinking"); + assert_eq!(p2["event_type"], "tool_started"); + assert_eq!(p3["event_type"], "tool_completed"); + assert_eq!(p4["event_type"], "response"); + + ws.close(None).await.unwrap(); +}