diff --git a/web-gateway/src/grpc/mod.rs b/web-gateway/src/grpc/mod.rs index 705f46dba..e6af67a86 100644 --- a/web-gateway/src/grpc/mod.rs +++ b/web-gateway/src/grpc/mod.rs @@ -1 +1,2 @@ pub mod clients; +pub mod streams; diff --git a/web-gateway/src/grpc/streams.rs b/web-gateway/src/grpc/streams.rs new file mode 100644 index 000000000..6ef5c7864 --- /dev/null +++ b/web-gateway/src/grpc/streams.rs @@ -0,0 +1,109 @@ +use std::time::Duration; + +use tokio::sync::broadcast; +use tonic::transport::Channel; + +use crate::ws::messages::ServerMessage; + +/// Start background tasks that bridge gRPC streaming RPCs to the WebSocket broadcast channel. +/// +/// Each stream task subscribes to a gRPC server-streaming RPC, converts events to +/// `ServerMessage` JSON, and sends them to the broadcast channel. Tasks reconnect +/// on failure with exponential backoff. +pub fn start_grpc_stream_bridges( + trading_channel: Option, + ws_broadcast: broadcast::Sender, +) { + if let Some(channel) = trading_channel { + // Market data stream bridge + let tx = ws_broadcast.clone(); + let ch = channel.clone(); + tokio::spawn(async move { + stream_with_reconnect("market_data", ch, tx, |_channel| async { + // TODO: Wire to SubscribeMarketData RPC when proto streaming is defined. + // For now, this is a placeholder that yields pending forever. + std::future::pending::<()>().await; + }) + .await; + }); + + // Order updates stream bridge + let tx = ws_broadcast.clone(); + let ch = channel.clone(); + tokio::spawn(async move { + stream_with_reconnect("order_update", ch, tx, |_channel| async { + std::future::pending::<()>().await; + }) + .await; + }); + + // Risk alerts stream bridge + let tx = ws_broadcast.clone(); + let ch = channel.clone(); + tokio::spawn(async move { + stream_with_reconnect("risk_alert", ch, tx, |_channel| async { + std::future::pending::<()>().await; + }) + .await; + }); + + // Metrics stream bridge + let tx = ws_broadcast.clone(); + let ch = channel.clone(); + tokio::spawn(async move { + stream_with_reconnect("metrics", ch, tx, |_channel| async { + std::future::pending::<()>().await; + }) + .await; + }); + } + + // Allow publishing synthetic test events for development + let tx = ws_broadcast; + tokio::spawn(async move { + // Periodically broadcast a heartbeat metrics message so WebSocket clients + // can verify connectivity (every 30 seconds) + let mut interval = tokio::time::interval(Duration::from_secs(30)); + loop { + interval.tick().await; + let msg = ServerMessage::Metrics { + data: serde_json::json!({ + "heartbeat": true, + "timestamp": chrono::Utc::now().to_rfc3339(), + }), + }; + if let Ok(json) = serde_json::to_string(&msg) { + // Ignore send errors (no subscribers) + let _ = tx.send(json); + } + } + }); +} + +/// Reconnect wrapper with exponential backoff for a gRPC stream bridge task. +async fn stream_with_reconnect( + stream_name: &str, + channel: Channel, + _tx: broadcast::Sender, + connect_fn: F, +) where + F: Fn(Channel) -> Fut, + Fut: std::future::Future, +{ + let mut backoff = Duration::from_secs(1); + let max_backoff = Duration::from_secs(60); + + loop { + tracing::info!("Connecting gRPC stream: {}", stream_name); + connect_fn(channel.clone()).await; + + // Stream ended — reconnect with backoff + tracing::warn!( + "gRPC stream {} disconnected, reconnecting in {:?}", + stream_name, + backoff + ); + tokio::time::sleep(backoff).await; + backoff = (backoff * 2).min(max_backoff); + } +} diff --git a/web-gateway/src/lib.rs b/web-gateway/src/lib.rs index cc65d990b..bfce99aac 100644 --- a/web-gateway/src/lib.rs +++ b/web-gateway/src/lib.rs @@ -4,6 +4,7 @@ pub mod error; pub mod grpc; pub mod routes; pub mod state; +pub mod ws; pub mod proto { pub mod trading { diff --git a/web-gateway/src/main.rs b/web-gateway/src/main.rs index 7e423b397..f2e203961 100644 --- a/web-gateway/src/main.rs +++ b/web-gateway/src/main.rs @@ -4,6 +4,7 @@ use tower_http::trace::TraceLayer; use tracing::info; use web_gateway::config::GatewayConfig; +use web_gateway::grpc::streams::start_grpc_stream_bridges; use web_gateway::routes::create_router; use web_gateway::state::AppState; @@ -21,6 +22,9 @@ async fn main() -> Result<()> { let state = AppState::new(config).await?; + // Start gRPC stream bridge tasks (forward gRPC streams to WebSocket broadcast) + start_grpc_stream_bridges(state.trading_channel.clone(), state.ws_broadcast.clone()); + let cors = CorsLayer::new() .allow_origin(Any) .allow_methods(Any) diff --git a/web-gateway/src/routes/mod.rs b/web-gateway/src/routes/mod.rs index 42bbcfc99..1044754df 100644 --- a/web-gateway/src/routes/mod.rs +++ b/web-gateway/src/routes/mod.rs @@ -1,7 +1,8 @@ -use axum::{middleware, Router}; +use axum::{middleware, routing, Router}; use crate::auth::middleware::auth_middleware; use crate::state::AppState; +use crate::ws::handler::ws_handler; pub mod backtesting; pub mod config; @@ -13,7 +14,8 @@ pub mod training; pub mod tune; pub fn create_router(state: AppState) -> Router { - let api = Router::new() + // REST endpoints require auth middleware + let rest_api = Router::new() .nest("/trading", trading::router()) .nest("/risk", risk::router()) .nest("/ml", ml::router()) @@ -27,5 +29,11 @@ pub fn create_router(state: AppState) -> Router { auth_middleware, )); - Router::new().nest("/api", api).with_state(state) + // WebSocket endpoint validates JWT via query param (no middleware layer) + let ws_route = Router::new().route("/ws", routing::get(ws_handler)); + + Router::new() + .nest("/api", rest_api) + .nest("/api", ws_route) + .with_state(state) } diff --git a/web-gateway/src/ws/handler.rs b/web-gateway/src/ws/handler.rs new file mode 100644 index 000000000..60c175b4d --- /dev/null +++ b/web-gateway/src/ws/handler.rs @@ -0,0 +1,126 @@ +use std::collections::HashSet; +use std::sync::Arc; + +use axum::{ + extract::{ + ws::{Message, WebSocket}, + Query, State, WebSocketUpgrade, + }, + response::IntoResponse, +}; +use futures_util::{SinkExt, StreamExt}; +use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; +use serde::Deserialize; + +use crate::auth::claims::Claims; +use crate::config::GatewayConfig; +use crate::error::AppError; +use crate::state::AppState; +use crate::ws::messages::{ClientMessage, ServerMessage}; + +/// Query parameters for WebSocket upgrade (JWT passed via query string) +#[derive(Debug, Deserialize)] +pub struct WsParams { + token: String, +} + +/// WebSocket upgrade handler — validates JWT from query param then upgrades +pub async fn ws_handler( + ws: WebSocketUpgrade, + State(state): State, + Query(params): Query, +) -> Result { + // Validate JWT from query parameter + let _claims = validate_ws_token(¶ms.token, &state.config)?; + + Ok(ws.on_upgrade(move |socket| handle_ws_connection(socket, state))) +} + +/// Validate JWT token from WebSocket query parameter +fn validate_ws_token(token: &str, config: &Arc) -> Result { + let validation = Validation::new(Algorithm::HS256); + let key = DecodingKey::from_secret(config.jwt_secret.as_bytes()); + + let token_data = decode::(token, &key, &validation).map_err(|_| AppError::Unauthorized)?; + + Ok(token_data.claims) +} + +/// Handle an established WebSocket connection +async fn handle_ws_connection(socket: WebSocket, state: AppState) { + let (mut ws_sender, mut ws_receiver) = socket.split(); + + // Subscribe to the broadcast channel for server events + let mut broadcast_rx = state.ws_broadcast.subscribe(); + + // Client's subscribed topics (empty = receive nothing until they subscribe) + let subscribed_topics: Arc>> = + Arc::new(tokio::sync::Mutex::new(HashSet::new())); + + let topics_for_broadcast = subscribed_topics.clone(); + + // Task: forward broadcast events to WebSocket client (filtered by subscribed topics) + let mut send_task = tokio::spawn(async move { + while let Ok(msg_json) = broadcast_rx.recv().await { + // Parse the broadcast message to check its topic + let should_send = if let Ok(server_msg) = + serde_json::from_str::(&msg_json) + { + let topics = topics_for_broadcast.lock().await; + topics.contains(server_msg.topic()) + || topics.contains("*") // wildcard subscription + } else { + false + }; + + if should_send { + if ws_sender.send(Message::Text(msg_json.into())).await.is_err() { + break; // Client disconnected + } + } + } + }); + + // Task: receive client messages (subscribe/unsubscribe) + let topics_for_recv = subscribed_topics.clone(); + let mut recv_task = tokio::spawn(async move { + while let Some(Ok(msg)) = ws_receiver.next().await { + match msg { + Message::Text(text) => { + if let Ok(client_msg) = serde_json::from_str::(&text) { + let mut topics = topics_for_recv.lock().await; + match client_msg { + ClientMessage::Subscribe { topics: new_topics } => { + for topic in new_topics { + topics.insert(topic); + } + } + ClientMessage::Unsubscribe { + topics: remove_topics, + } => { + for topic in &remove_topics { + topics.remove(topic); + } + } + } + } + // Silently ignore malformed messages + } + Message::Close(_) => break, + _ => {} // Ignore ping/pong/binary + } + } + }); + + // Wait for either task to finish, then abort the other + tokio::select! { + _ = &mut send_task => { + recv_task.abort(); + } + _ = &mut recv_task => { + send_task.abort(); + } + } + + tracing::debug!("WebSocket connection closed"); +} diff --git a/web-gateway/src/ws/messages.rs b/web-gateway/src/ws/messages.rs new file mode 100644 index 000000000..78ce3bf6e --- /dev/null +++ b/web-gateway/src/ws/messages.rs @@ -0,0 +1,65 @@ +use serde::{Deserialize, Serialize}; + +/// Client -> Server messages sent over WebSocket +#[derive(Debug, Deserialize)] +#[serde(tag = "type")] +pub enum ClientMessage { + /// Subscribe to one or more data topics + #[serde(rename = "subscribe")] + Subscribe { topics: Vec }, + /// Unsubscribe from one or more data topics + #[serde(rename = "unsubscribe")] + Unsubscribe { topics: Vec }, +} + +/// Server -> Client messages broadcast over WebSocket +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum ServerMessage { + /// Real-time market data update + #[serde(rename = "market_data")] + MarketData { + symbol: String, + data: serde_json::Value, + }, + /// Order status change + #[serde(rename = "order_update")] + OrderUpdate { data: serde_json::Value }, + /// Risk threshold alert + #[serde(rename = "risk_alert")] + RiskAlert { + severity: String, + data: serde_json::Value, + }, + /// Position change notification + #[serde(rename = "position_update")] + PositionUpdate { data: serde_json::Value }, + /// ML model prediction + #[serde(rename = "ml_prediction")] + MlPrediction { data: serde_json::Value }, + /// Training job progress + #[serde(rename = "training_progress")] + TrainingProgress { data: serde_json::Value }, + /// System metrics snapshot + #[serde(rename = "metrics")] + Metrics { data: serde_json::Value }, + /// Configuration change notification + #[serde(rename = "config_update")] + ConfigUpdate { data: serde_json::Value }, +} + +impl ServerMessage { + /// Extract the topic string for this message (used for client-side filtering) + pub fn topic(&self) -> &str { + match self { + ServerMessage::MarketData { .. } => "market_data", + ServerMessage::OrderUpdate { .. } => "order_update", + ServerMessage::RiskAlert { .. } => "risk_alert", + ServerMessage::PositionUpdate { .. } => "position_update", + ServerMessage::MlPrediction { .. } => "ml_prediction", + ServerMessage::TrainingProgress { .. } => "training_progress", + ServerMessage::Metrics { .. } => "metrics", + ServerMessage::ConfigUpdate { .. } => "config_update", + } + } +} diff --git a/web-gateway/src/ws/mod.rs b/web-gateway/src/ws/mod.rs new file mode 100644 index 000000000..b8661994a --- /dev/null +++ b/web-gateway/src/ws/mod.rs @@ -0,0 +1,2 @@ +pub mod handler; +pub mod messages;