feat(web-gateway): WebSocket infrastructure with topic-based subscriptions

Add WebSocket support:
- ws/messages.rs: ClientMessage (subscribe/unsubscribe) and ServerMessage
  (market_data, order_update, risk_alert, position_update, ml_prediction,
  training_progress, metrics, config_update)
- ws/handler.rs: WebSocket upgrade with JWT validation from query param,
  per-client topic filtering via HashSet, dual-task architecture
  (send from broadcast, receive client commands)
- grpc/streams.rs: Background bridge tasks connecting gRPC streaming RPCs
  to broadcast channel with exponential backoff reconnection.
  Includes 30s heartbeat for connectivity verification.

WebSocket endpoint at GET /api/ws?token=<jwt>

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-21 23:54:09 +01:00
parent e364d447f5
commit a1affb0767
8 changed files with 319 additions and 3 deletions

View File

@@ -1 +1,2 @@
pub mod clients;
pub mod streams;

View File

@@ -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<Channel>,
ws_broadcast: broadcast::Sender<String>,
) {
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<F, Fut>(
stream_name: &str,
channel: Channel,
_tx: broadcast::Sender<String>,
connect_fn: F,
) where
F: Fn(Channel) -> Fut,
Fut: std::future::Future<Output = ()>,
{
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);
}
}

View File

@@ -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 {

View File

@@ -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)

View File

@@ -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)
}

View File

@@ -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<AppState>,
Query(params): Query<WsParams>,
) -> Result<impl IntoResponse, AppError> {
// Validate JWT from query parameter
let _claims = validate_ws_token(&params.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<GatewayConfig>) -> Result<Claims, AppError> {
let validation = Validation::new(Algorithm::HS256);
let key = DecodingKey::from_secret(config.jwt_secret.as_bytes());
let token_data = decode::<Claims>(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<tokio::sync::Mutex<HashSet<String>>> =
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::<ServerMessage>(&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::<ClientMessage>(&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");
}

View File

@@ -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<String> },
/// Unsubscribe from one or more data topics
#[serde(rename = "unsubscribe")]
Unsubscribe { topics: Vec<String> },
}
/// 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",
}
}
}

View File

@@ -0,0 +1,2 @@
pub mod handler;
pub mod messages;