Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 74 additions & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 2 additions & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"] }

Expand Down Expand Up @@ -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"
Expand Down
9 changes: 7 additions & 2 deletions src/channels/web/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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;
Expand Down Expand Up @@ -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 {
Expand All @@ -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);
Expand Down Expand Up @@ -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)))
Expand Down
53 changes: 50 additions & 3 deletions src/channels/web/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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::{
Expand Down Expand Up @@ -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));
Expand All @@ -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))
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -137,7 +151,7 @@ pub async fn start_server(
}
});

Ok(())
Ok(bound_addr)
}

// --- Static file handlers ---
Expand Down Expand Up @@ -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>,
Expand Down Expand Up @@ -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,
}
66 changes: 66 additions & 0 deletions src/channels/web/sse.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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);
}
}
Loading