mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-26 15:40:18 +00:00
Add WebSocket gateway and control plane (#8)
* Add WebSocket gateway and control plane endpoint Adds bidirectional WebSocket transport to the web gateway alongside the existing SSE stream. Clients can send messages, approvals, and pings over a single persistent connection at /api/chat/ws. - Enable axum `ws` feature for built-in WebSocket support - Add WsClientMessage/WsServerMessage types with tagged JSON protocol - Add subscribe_raw() to SseManager for non-SSE consumers - Create ws.rs with connection handler (split sender/receiver tasks) - Add WsConnectionTracker for active connection counting - Add /api/gateway/status control plane endpoint (SSE + WS counts) - 35 new tests covering message types, broadcast, and handler logic https://claude.ai/code/session_01KEaLN6Xq2j5EeV3SGHQT6b * Add e2e WebSocket gateway integration tests - Add tokio-tungstenite dev-dependency for WebSocket client in tests - Update start_server to return actual bound SocketAddr (enables port 0) - Add 10 e2e tests covering full HTTP upgrade → WebSocket → message flow: ping/pong, message routing to agent, broadcast event delivery, connection tracking, invalid message handling, auth rejection, gateway status endpoint, and multi-event sequencing https://claude.ai/code/session_01KEaLN6Xq2j5EeV3SGHQT6b --------- Co-authored-by: Claude <[email protected]>
This commit is contained in:
Generated
+74
@@ -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"
|
||||
|
||||
+2
-1
@@ -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"
|
||||
|
||||
@@ -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)))
|
||||
|
||||
@@ -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<Option<oneshot::Sender<()>>>,
|
||||
/// WebSocket connection tracker.
|
||||
pub ws_tracker: Option<Arc<crate::channels::web::ws::WsConnectionTracker>>,
|
||||
}
|
||||
|
||||
/// 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<GatewayState>,
|
||||
auth_token: String,
|
||||
) -> Result<(), crate::error::ChannelError> {
|
||||
) -> Result<SocketAddr, crate::error::ChannelError> {
|
||||
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<Arc<GatewayState>>) -> impl Int
|
||||
state.sse.subscribe()
|
||||
}
|
||||
|
||||
async fn chat_ws_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
) -> impl IntoResponse {
|
||||
ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct HistoryQuery {
|
||||
thread_id: Option<String>,
|
||||
@@ -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<Arc<GatewayState>>,
|
||||
) -> Json<GatewayStatusResponse> {
|
||||
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,
|
||||
}
|
||||
|
||||
@@ -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<Item = SseEvent> + 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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<String>,
|
||||
},
|
||||
/// 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<WsClientMessage, _> = 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"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<GatewayState>) {
|
||||
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::<WsServerMessage>(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<WsClientMessage, _> = 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<WsServerMessage>,
|
||||
) {
|
||||
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<mpsc::Sender<IncomingMessage>>) -> 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())),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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<GatewayState>,
|
||||
mpsc::Receiver<IncomingMessage>,
|
||||
) {
|
||||
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<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>> {
|
||||
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<Item = Result<Message, tokio_tungstenite::tungstenite::Error>> + 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();
|
||||
}
|
||||
Reference in New Issue
Block a user