#![allow(clippy::panic)] use std::{collections::HashSet, net::SocketAddr, sync::Arc, time::Duration};
use axum::{Router, routing::get};
use fraiseql_core::{runtime::SubscriptionManager, schema::CompiledSchema};
use futures::{SinkExt, StreamExt, future::BoxFuture};
use tokio::{net::TcpListener, sync::mpsc};
use tokio_tungstenite::{
connect_async,
tungstenite::{self, client::IntoClientRequest},
};
use super::{
context_hash::security_context_hash,
delivery::{EntityEvent, EventDeliveryPipeline, EventKindSerde, RlsEvaluator},
observer::RealtimeBroadcastObserver,
routes::{RealtimeSchemaConfig, realtime_router},
server::{
RealtimeConfig, RealtimeServer, RealtimeState, TokenInfo, TokenValidator, ws_handler,
},
};
#[derive(Clone)]
struct TestValidator {
expires_in: i64,
}
impl TestValidator {
const fn new() -> Self {
Self { expires_in: 3600 }
}
const fn with_expires_in(expires_in: i64) -> Self {
Self { expires_in }
}
}
impl TokenValidator for TestValidator {
fn validate<'a>(&'a self, token: &'a str) -> BoxFuture<'a, Result<TokenInfo, String>> {
let expires_in = self.expires_in;
let token = token.to_owned();
Box::pin(async move {
if token.starts_with("valid-") {
let user_id = token.strip_prefix("valid-").unwrap_or("unknown").to_owned();
Ok(TokenInfo {
user_id: user_id.clone(),
context_hash: {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
user_id.hash(&mut hasher);
hasher.finish()
},
expires_at: chrono::Utc::now().timestamp() + expires_in,
})
} else if token == "expired-token" {
Err("token expired".to_owned())
} else {
Err("invalid token".to_owned())
}
})
}
}
async fn spawn_test_server(config: RealtimeConfig, validator: TestValidator) -> SocketAddr {
let server = Arc::new(RealtimeServer::new(config));
let state = RealtimeState {
server,
validator: Arc::new(validator),
};
let app = Router::new().route("/realtime/v1", get(ws_handler)).with_state(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
addr
}
fn ws_url(addr: SocketAddr, token: Option<&str>) -> String {
match token {
Some(t) => format!("ws://{addr}/realtime/v1?token={t}"),
None => format!("ws://{addr}/realtime/v1"),
}
}
async fn spawn_test_server_with_entities(
config: RealtimeConfig,
validator: TestValidator,
entities: HashSet<String>,
) -> SocketAddr {
let server = Arc::new(RealtimeServer::with_entities(config, entities));
let state = RealtimeState {
server,
validator: Arc::new(validator),
};
let app = Router::new().route("/realtime/v1", get(ws_handler)).with_state(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
addr
}
fn parse_server_msg(msg: &tungstenite::Message) -> serde_json::Value {
match msg {
tungstenite::Message::Text(text) => serde_json::from_str(text).unwrap(),
other => panic!("Expected text message, got {other:?}"),
}
}
async fn send_json(
ws: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
msg: serde_json::Value,
) {
ws.send(tungstenite::Message::Text(msg.to_string().into())).await.unwrap();
}
async fn next_msg(
ws: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> serde_json::Value {
let msg = tokio::time::timeout(Duration::from_secs(5), ws.next())
.await
.expect("timed out waiting for message")
.expect("stream ended")
.expect("WebSocket error");
parse_server_msg(&msg)
}
#[tokio::test]
async fn test_websocket_connect_with_valid_token() {
let addr = spawn_test_server(RealtimeConfig::default(), TestValidator::new()).await;
let (mut ws, _response) = connect_async(ws_url(addr, Some("valid-alice"))).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
assert!(parsed["connection_id"].is_string());
assert!(!parsed["connection_id"].as_str().unwrap().is_empty());
ws.close(None).await.unwrap();
}
#[tokio::test]
async fn test_websocket_connect_without_token_returns_401() {
let addr = spawn_test_server(RealtimeConfig::default(), TestValidator::new()).await;
let result = connect_async(ws_url(addr, None)).await;
assert!(result.is_err(), "Expected connection to be rejected");
if let Err(tungstenite::Error::Http(response)) = result {
assert_eq!(response.status(), 401);
} else {
panic!("Expected HTTP error with 401 status");
}
}
#[tokio::test]
async fn test_websocket_connect_with_expired_token_returns_401() {
let addr = spawn_test_server(RealtimeConfig::default(), TestValidator::new()).await;
let result = connect_async(ws_url(addr, Some("expired-token"))).await;
assert!(result.is_err(), "Expected connection to be rejected");
if let Err(tungstenite::Error::Http(response)) = result {
assert_eq!(response.status(), 401);
} else {
panic!("Expected HTTP error with 401 status");
}
}
#[tokio::test]
async fn test_websocket_heartbeat_pong() {
let config = RealtimeConfig {
heartbeat_interval: Duration::from_millis(100),
idle_timeout: Duration::from_secs(10),
..RealtimeConfig::default()
};
let addr = spawn_test_server(config, TestValidator::new()).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-bob"))).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "ping");
let pong = serde_json::json!({"type": "pong"}).to_string();
ws.send(tungstenite::Message::Text(pong.into())).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "ping");
ws.close(None).await.unwrap();
}
#[tokio::test]
async fn test_websocket_idle_timeout_disconnects() {
let config = RealtimeConfig {
heartbeat_interval: Duration::from_secs(60),
idle_timeout: Duration::from_millis(200),
..RealtimeConfig::default()
};
let addr = spawn_test_server(config, TestValidator::new()).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-carol"))).await.unwrap();
let _connected = ws.next().await.unwrap().unwrap();
let msg = ws.next().await;
match msg {
Some(Ok(tungstenite::Message::Close(Some(frame)))) => {
assert_eq!(frame.code, tungstenite::protocol::frame::coding::CloseCode::Normal);
},
None => {},
other => {
panic!("Expected close frame or connection drop, got {other:?}");
},
}
}
#[tokio::test]
async fn test_websocket_graceful_close() {
let addr = spawn_test_server(RealtimeConfig::default(), TestValidator::new()).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-dave"))).await.unwrap();
let _connected = ws.next().await.unwrap().unwrap();
ws.close(None).await.unwrap();
let msg = ws.next().await;
let is_clean_close = matches!(
msg,
None | Some(
Ok(tungstenite::Message::Close(_))
| Err(tungstenite::Error::Protocol(
tungstenite::error::ProtocolError::ResetWithoutClosingHandshake
) | tungstenite::Error::ConnectionClosed)
)
);
assert!(is_clean_close, "Expected close or None, got {msg:?}");
}
#[tokio::test]
async fn test_websocket_connection_limit_per_context() {
let config = RealtimeConfig {
max_connections_per_context: 2,
..RealtimeConfig::default()
};
let addr = spawn_test_server(config, TestValidator::new()).await;
let (mut ws1, _) = connect_async(ws_url(addr, Some("valid-eve"))).await.unwrap();
let _connected1 = ws1.next().await.unwrap().unwrap();
let (mut ws2, _) = connect_async(ws_url(addr, Some("valid-eve"))).await.unwrap();
let _connected2 = ws2.next().await.unwrap().unwrap();
let result = connect_async(ws_url(addr, Some("valid-eve"))).await;
assert!(result.is_err(), "Expected third connection to be rejected");
if let Err(tungstenite::Error::Http(response)) = result {
assert_eq!(response.status(), 429);
} else {
panic!("Expected HTTP 429 error");
}
let (mut ws3, _) = connect_async(ws_url(addr, Some("valid-frank"))).await.unwrap();
let msg = ws3.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
ws1.close(None).await.ok();
ws2.close(None).await.ok();
ws3.close(None).await.ok();
}
#[tokio::test]
async fn test_websocket_token_expiry_disconnects() {
let config = RealtimeConfig {
heartbeat_interval: Duration::from_millis(100),
idle_timeout: Duration::from_secs(10),
..RealtimeConfig::default()
};
let validator = TestValidator::with_expires_in(1);
let addr = spawn_test_server(config, validator).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-grace"))).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
let start = std::time::Instant::now();
let mut got_token_expired = false;
while start.elapsed() < Duration::from_secs(3) {
match ws.next().await {
Some(Ok(tungstenite::Message::Text(text))) => {
let parsed: serde_json::Value = serde_json::from_str(&text).unwrap();
if parsed["type"] == "token_expired" {
got_token_expired = true;
break;
}
assert_eq!(parsed["type"], "ping");
},
Some(Ok(tungstenite::Message::Close(Some(frame)))) => {
assert_eq!(frame.code, tungstenite::protocol::frame::coding::CloseCode::from(4401));
got_token_expired = true;
break;
},
Some(Ok(tungstenite::Message::Close(None))) | None => break,
other => panic!("Unexpected message: {other:?}"),
}
}
assert!(got_token_expired, "Expected token_expired message before close");
}
#[derive(Clone)]
struct NearExpiryValidator;
impl TokenValidator for NearExpiryValidator {
fn validate<'a>(&'a self, token: &'a str) -> BoxFuture<'a, Result<TokenInfo, String>> {
let token = token.to_owned();
Box::pin(async move {
if token.starts_with("valid-") {
let user_id = token.strip_prefix("valid-").unwrap_or("unknown").to_owned();
Ok(TokenInfo {
user_id,
context_hash: 42,
expires_at: chrono::Utc::now().timestamp(),
})
} else {
Err("invalid".to_owned())
}
})
}
}
#[tokio::test]
async fn test_websocket_token_revalidation_interval() {
let config = RealtimeConfig {
heartbeat_interval: Duration::from_millis(50),
idle_timeout: Duration::from_secs(10),
token_revalidation_interval: Duration::from_millis(50),
..RealtimeConfig::default()
};
let server = Arc::new(RealtimeServer::new(config));
let state = RealtimeState {
server,
validator: Arc::new(NearExpiryValidator),
};
let app = Router::new().route("/realtime/v1", get(ws_handler)).with_state(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-henry"))).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
let mut got_expired = false;
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
while tokio::time::Instant::now() < deadline {
match tokio::time::timeout(Duration::from_millis(500), ws.next()).await {
Ok(Some(Ok(tungstenite::Message::Text(text)))) => {
let parsed: serde_json::Value = serde_json::from_str(&text).unwrap();
if parsed["type"] == "token_expired" {
got_expired = true;
break;
}
},
Ok(Some(Ok(tungstenite::Message::Close(Some(frame))))) => {
if frame.code == tungstenite::protocol::frame::coding::CloseCode::from(4401) {
got_expired = true;
}
break;
},
Ok(None | Some(Ok(tungstenite::Message::Close(None)))) => break,
_ => {},
}
}
assert!(got_expired, "Expected token_expired from revalidation");
}
fn test_entities() -> HashSet<String> {
["Post", "Comment"].iter().map(|s| (*s).to_owned()).collect()
}
#[tokio::test]
async fn test_subscribe_to_entity() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-alice"))).await.unwrap();
let connected = next_msg(&mut ws).await;
assert_eq!(connected["type"], "connected");
send_json(
&mut ws,
serde_json::json!({"type": "subscribe", "entity": "Post", "event": "*"}),
)
.await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Post");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_subscribe_with_event_filter() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-bob"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(
&mut ws,
serde_json::json!({"type": "subscribe", "entity": "Post", "event": "INSERT"}),
)
.await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Post");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_subscribe_with_field_filter() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-carol"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(
&mut ws,
serde_json::json!({"type": "subscribe", "entity": "Post", "filter": "author_id=eq.123"}),
)
.await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Post");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_unsubscribe() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-dave"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
send_json(&mut ws, serde_json::json!({"type": "unsubscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "unsubscribed");
assert_eq!(reply["entity"], "Post");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_subscribe_to_nonexistent_entity_returns_error() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-eve"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Foo"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "error");
assert!(
reply["message"].as_str().unwrap().contains("unknown entity"),
"Expected 'unknown entity' error, got: {}",
reply["message"]
);
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_subscribe_exceeds_fan_out_limit() {
let config = RealtimeConfig {
max_subscriptions_per_entity: 2,
max_connections_per_context: 100,
..RealtimeConfig::default()
};
let addr = spawn_test_server_with_entities(config, TestValidator::new(), test_entities()).await;
let (mut ws1, _) = connect_async(ws_url(addr, Some("valid-user1"))).await.unwrap();
let _ = next_msg(&mut ws1).await;
send_json(&mut ws1, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws1).await;
assert_eq!(reply["type"], "subscribed");
let (mut ws2, _) = connect_async(ws_url(addr, Some("valid-user2"))).await.unwrap();
let _ = next_msg(&mut ws2).await;
send_json(&mut ws2, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws2).await;
assert_eq!(reply["type"], "subscribed");
let (mut ws3, _) = connect_async(ws_url(addr, Some("valid-user3"))).await.unwrap();
let _ = next_msg(&mut ws3).await;
send_json(&mut ws3, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws3).await;
assert_eq!(reply["type"], "error");
assert!(
reply["message"].as_str().unwrap().contains("limit"),
"Expected fan-out limit error, got: {}",
reply["message"]
);
ws1.close(None).await.ok();
ws2.close(None).await.ok();
ws3.close(None).await.ok();
}
#[tokio::test]
async fn test_multiple_subscriptions_same_client() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-frank"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Post");
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Comment"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Comment");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_duplicate_subscribe_is_idempotent() {
let addr = spawn_test_server_with_entities(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-grace"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
assert_eq!(reply["entity"], "Post");
ws.close(None).await.ok();
}
struct AllowAllRls;
impl RlsEvaluator for AllowAllRls {
fn can_access<'a>(
&'a self,
_context_hash: u64,
_entity: &'a str,
_row: &'a serde_json::Value,
) -> BoxFuture<'a, bool> {
Box::pin(async { true })
}
}
struct DenyAllRls;
impl RlsEvaluator for DenyAllRls {
fn can_access<'a>(
&'a self,
_context_hash: u64,
_entity: &'a str,
_row: &'a serde_json::Value,
) -> BoxFuture<'a, bool> {
Box::pin(async { false })
}
}
struct CountingRls {
call_count: std::sync::atomic::AtomicUsize,
}
impl CountingRls {
fn new() -> Self {
Self {
call_count: std::sync::atomic::AtomicUsize::new(0),
}
}
fn count(&self) -> usize {
self.call_count.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl RlsEvaluator for CountingRls {
fn can_access<'a>(
&'a self,
_context_hash: u64,
_entity: &'a str,
_row: &'a serde_json::Value,
) -> BoxFuture<'a, bool> {
self.call_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Box::pin(async { true })
}
}
async fn spawn_server_with_delivery(
config: RealtimeConfig,
validator: TestValidator,
entities: HashSet<String>,
rls: Arc<dyn RlsEvaluator>,
) -> (SocketAddr, mpsc::Sender<EntityEvent>) {
let server = Arc::new(RealtimeServer::with_entities(config, entities));
let (event_tx, event_rx) = mpsc::channel(1000);
let pipeline = EventDeliveryPipeline::new(
server.subscriptions.clone(),
server.connections.clone(),
rls,
event_rx,
);
tokio::spawn(pipeline.run());
let state = RealtimeState {
server,
validator: Arc::new(validator),
};
let app = Router::new().route("/realtime/v1", get(ws_handler)).with_state(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
(addr, event_tx)
}
fn make_post_event(event_kind: EventKindSerde, author_id: i64) -> EntityEvent {
EntityEvent {
entity: "Post".to_owned(),
event_kind,
new: Some(serde_json::json!({"id": 1, "author_id": author_id, "title": "Hello"})),
old: None,
timestamp: "2026-04-28T12:00:00Z".to_owned(),
}
}
#[tokio::test]
async fn test_event_delivered_to_subscribed_client() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-alice"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed");
event_tx.send(make_post_event(EventKindSerde::Insert, 42)).await.unwrap();
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["entity"], "Post");
assert_eq!(msg["event"], "INSERT");
assert_eq!(msg["new"]["author_id"], 42);
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_not_delivered_to_unsubscribed_client() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-bob"))).await.unwrap();
let _ = next_msg(&mut ws).await;
event_tx.send(make_post_event(EventKindSerde::Insert, 42)).await.unwrap();
let result = tokio::time::timeout(Duration::from_millis(200), ws.next()).await;
assert!(result.is_err(), "Expected timeout (no message), got a message");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_rls_filters_unauthorized() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(DenyAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-carol"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
event_tx.send(make_post_event(EventKindSerde::Insert, 42)).await.unwrap();
let result = tokio::time::timeout(Duration::from_millis(200), ws.next()).await;
assert!(result.is_err(), "Expected no message (RLS denied), got a message");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_rls_allows_authorized() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-dave"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
event_tx.send(make_post_event(EventKindSerde::Update, 99)).await.unwrap();
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["event"], "UPDATE");
assert_eq!(msg["new"]["author_id"], 99);
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_rls_grouping_by_context_hash() {
let counting_rls = Arc::new(CountingRls::new());
let config = RealtimeConfig {
max_connections_per_context: 100,
..RealtimeConfig::default()
};
let (addr, event_tx) = spawn_server_with_delivery(
config,
TestValidator::new(),
test_entities(),
counting_rls.clone(),
)
.await;
let mut same_user_ws = Vec::new();
for i in 0..3 {
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-sameuser"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let reply = next_msg(&mut ws).await;
assert_eq!(reply["type"], "subscribed", "client {i} failed to subscribe");
same_user_ws.push(ws);
}
let (mut ws_diff1, _) = connect_async(ws_url(addr, Some("valid-other1"))).await.unwrap();
let _ = next_msg(&mut ws_diff1).await;
send_json(&mut ws_diff1, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws_diff1).await;
let (mut ws_diff2, _) = connect_async(ws_url(addr, Some("valid-other2"))).await.unwrap();
let _ = next_msg(&mut ws_diff2).await;
send_json(&mut ws_diff2, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws_diff2).await;
event_tx.send(make_post_event(EventKindSerde::Insert, 1)).await.unwrap();
for ws in &mut same_user_ws {
let msg = next_msg(ws).await;
assert_eq!(msg["type"], "change");
}
let msg = next_msg(&mut ws_diff1).await;
assert_eq!(msg["type"], "change");
let msg = next_msg(&mut ws_diff2).await;
assert_eq!(msg["type"], "change");
let rls_calls = counting_rls.count();
assert_eq!(
rls_calls, 3,
"Expected 3 RLS evaluations (one per distinct context hash), got {rls_calls}"
);
for ws in &mut same_user_ws {
ws.close(None).await.ok();
}
ws_diff1.close(None).await.ok();
ws_diff2.close(None).await.ok();
}
#[tokio::test]
async fn test_event_field_filter_applied() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-eve"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(
&mut ws,
serde_json::json!({"type": "subscribe", "entity": "Post", "filter": "author_id=eq.123"}),
)
.await;
let _ = next_msg(&mut ws).await;
event_tx.send(make_post_event(EventKindSerde::Insert, 456)).await.unwrap();
let result = tokio::time::timeout(Duration::from_millis(200), ws.next()).await;
assert!(result.is_err(), "Expected no message for author_id=456");
event_tx.send(make_post_event(EventKindSerde::Insert, 123)).await.unwrap();
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["new"]["author_id"], 123);
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_type_filter_applied() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-frank"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(
&mut ws,
serde_json::json!({"type": "subscribe", "entity": "Post", "event": "INSERT"}),
)
.await;
let _ = next_msg(&mut ws).await;
event_tx.send(make_post_event(EventKindSerde::Update, 42)).await.unwrap();
let result = tokio::time::timeout(Duration::from_millis(200), ws.next()).await;
assert!(result.is_err(), "Expected no message for UPDATE event");
event_tx.send(make_post_event(EventKindSerde::Insert, 42)).await.unwrap();
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["event"], "INSERT");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_event_payload_format() {
let (addr, event_tx) = spawn_server_with_delivery(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-grace"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
let event = EntityEvent {
entity: "Post".to_owned(),
event_kind: EventKindSerde::Delete,
new: None,
old: Some(serde_json::json!({"id": 7, "title": "Deleted post"})),
timestamp: "2026-04-28T15:30:00Z".to_owned(),
};
event_tx.send(event).await.unwrap();
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["entity"], "Post");
assert_eq!(msg["event"], "DELETE");
assert!(msg["new"].is_null());
assert_eq!(msg["old"]["id"], 7);
assert_eq!(msg["old"]["title"], "Deleted post");
assert_eq!(msg["timestamp"], "2026-04-28T15:30:00Z");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_observer_receives_mutation_event() {
let (observer, mut event_rx) = RealtimeBroadcastObserver::new(100);
let event = make_post_event(EventKindSerde::Insert, 1);
observer.on_mutation_complete(event.clone());
let received = event_rx.recv().await.unwrap();
assert_eq!(received.entity, event.entity);
assert_eq!(received.event_kind, event.event_kind);
}
#[tokio::test]
async fn test_observer_enqueues_event_to_delivery_pipeline() {
let (observer, mut event_rx) = RealtimeBroadcastObserver::new(100);
for i in 0..5_i64 {
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, i));
}
let mut received = Vec::new();
for _ in 0..5 {
received.push(event_rx.try_recv().unwrap());
}
assert_eq!(received.len(), 5);
}
#[tokio::test]
async fn test_observer_returns_immediately() {
let (observer, _rx) = RealtimeBroadcastObserver::new(100);
let event = make_post_event(EventKindSerde::Insert, 1);
let start = std::time::Instant::now();
observer.on_mutation_complete(event);
let elapsed = start.elapsed();
assert!(
elapsed.as_millis() < 1,
"on_mutation_complete took {}µs — expected <1ms",
elapsed.as_micros()
);
}
#[tokio::test]
async fn test_observer_channel_full_drops_event() {
let (observer, _rx) = RealtimeBroadcastObserver::new(1);
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, 1));
assert_eq!(observer.events_dropped_total(), 0, "first send must succeed");
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, 2));
assert_eq!(observer.events_dropped_total(), 1, "second send must be dropped");
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, 3));
assert_eq!(observer.events_dropped_total(), 2, "third send must be dropped");
}
async fn spawn_server_with_observer(
config: RealtimeConfig,
validator: TestValidator,
entities: HashSet<String>,
rls: Arc<dyn RlsEvaluator>,
) -> (SocketAddr, RealtimeBroadcastObserver) {
let server = Arc::new(RealtimeServer::with_entities(config.clone(), entities));
let (observer, event_rx) = RealtimeBroadcastObserver::new(config.event_channel_capacity);
let pipeline = EventDeliveryPipeline::new(
server.subscriptions.clone(),
server.connections.clone(),
rls,
event_rx,
);
tokio::spawn(pipeline.run());
let state = RealtimeState {
server,
validator: Arc::new(validator),
};
let app = Router::new().route("/realtime/v1", get(ws_handler)).with_state(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
(addr, observer)
}
#[tokio::test]
async fn test_observer_slow_client_disconnected() {
let config = RealtimeConfig {
max_consecutive_drops: 3,
connection_event_capacity: 3,
max_connections_per_context: 100,
..RealtimeConfig::default()
};
let (addr, observer) = spawn_server_with_observer(
config,
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-slow"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, 1));
let _ = next_msg(&mut ws).await;
for i in 2..=8_i64 {
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, i));
}
let deadline = tokio::time::Instant::now() + Duration::from_secs(3);
let mut got_close_4002 = false;
while tokio::time::Instant::now() < deadline {
match tokio::time::timeout(Duration::from_millis(200), ws.next()).await {
Ok(Some(Ok(tungstenite::Message::Close(Some(frame))))) => {
if frame.code == tungstenite::protocol::frame::coding::CloseCode::from(4002) {
got_close_4002 = true;
}
break;
},
Ok(Some(Ok(_))) => {}, Ok(None | Some(Err(_))) | Err(_) => break, }
}
assert!(got_close_4002, "Expected close frame with code 4002 for slow consumer");
}
#[tokio::test]
async fn test_observer_slow_client_counter_resets() {
let config = RealtimeConfig {
max_consecutive_drops: 5,
connection_event_capacity: 10,
max_connections_per_context: 100,
..RealtimeConfig::default()
};
let (addr, observer) = spawn_server_with_observer(
config,
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-reader"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
for i in 0..10_i64 {
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, i));
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
}
observer.on_mutation_complete(make_post_event(EventKindSerde::Insert, 99));
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["new"]["author_id"], 99);
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_observer_end_to_end() {
let (addr, observer) = spawn_server_with_observer(
RealtimeConfig::default(),
TestValidator::new(),
test_entities(),
Arc::new(AllowAllRls),
)
.await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-e2e"))).await.unwrap();
let _ = next_msg(&mut ws).await;
send_json(&mut ws, serde_json::json!({"type": "subscribe", "entity": "Post"})).await;
let _ = next_msg(&mut ws).await;
observer.on_mutation_complete(EntityEvent {
entity: "Post".to_owned(),
event_kind: EventKindSerde::Insert,
new: Some(serde_json::json!({"id": 42, "title": "Hello realtime", "author_id": 7})),
old: None,
timestamp: "2026-04-29T00:00:00Z".to_owned(),
});
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "change");
assert_eq!(msg["entity"], "Post");
assert_eq!(msg["event"], "INSERT");
assert_eq!(msg["new"]["id"], 42);
assert_eq!(msg["new"]["title"], "Hello realtime");
assert!(msg["old"].is_null());
assert_eq!(msg["timestamp"], "2026-04-29T00:00:00Z");
ws.close(None).await.ok();
}
fn make_ctx<'a>(
user_id: &'a str,
roles: &'a [&'a str],
tenant_id: Option<&'a str>,
scopes: &'a [&'a str],
) -> super::context_hash::SecurityContextHashInput<'a> {
super::context_hash::SecurityContextHashInput {
user_id,
roles,
tenant_id,
scopes,
}
}
#[test]
fn test_security_context_hash_stable() {
let ctx = make_ctx("user-1", &["admin", "editor"], Some("tenant-A"), &["read:post"]);
let h1 = security_context_hash(&ctx);
let h2 = security_context_hash(&ctx);
assert_eq!(h1, h2, "same context must produce same hash");
}
#[test]
fn test_security_context_hash_differs_on_role_change() {
let ctx_a = make_ctx("user-1", &["admin"], None, &[]);
let ctx_b = make_ctx("user-1", &["editor"], None, &[]);
assert_ne!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"different roles must produce different hashes"
);
}
#[test]
fn test_security_context_hash_ignores_role_order() {
let ctx_a = make_ctx("user-1", &["admin", "editor"], None, &[]);
let ctx_b = make_ctx("user-1", &["editor", "admin"], None, &[]);
assert_eq!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"role order must not affect hash"
);
}
#[test]
fn test_security_context_hash_ignores_scope_order() {
let ctx_a = make_ctx("user-1", &[], None, &["read:post", "write:post"]);
let ctx_b = make_ctx("user-1", &[], None, &["write:post", "read:post"]);
assert_eq!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"scope order must not affect hash"
);
}
#[test]
fn test_security_context_hash_differs_on_user_id() {
let ctx_a = make_ctx("user-1", &["admin"], None, &[]);
let ctx_b = make_ctx("user-2", &["admin"], None, &[]);
assert_ne!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"different user IDs must produce different hashes"
);
}
#[test]
fn test_security_context_hash_differs_on_tenant() {
let ctx_a = make_ctx("user-1", &[], Some("tenant-A"), &[]);
let ctx_b = make_ctx("user-1", &[], Some("tenant-B"), &[]);
assert_ne!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"different tenant IDs must produce different hashes"
);
}
#[test]
fn test_security_context_hash_tenant_none_vs_some() {
let ctx_a = make_ctx("user-1", &[], None, &[]);
let ctx_b = make_ctx("user-1", &[], Some("tenant-A"), &[]);
assert_ne!(
security_context_hash(&ctx_a),
security_context_hash(&ctx_b),
"absent vs present tenant must produce different hashes"
);
}
#[test]
fn test_connection_state_stores_context_hash() {
use super::connections::ConnectionState;
let state = ConnectionState::new(
"conn-xyz".to_owned(),
"user-1".to_owned(),
0xdead_beef,
9_999_999_999,
);
assert_eq!(state.context_hash, 0xdead_beef);
assert_eq!(state.user_id, "user-1");
assert_eq!(state.connection_id, "conn-xyz");
}
#[tokio::test]
async fn test_connection_state_cleanup_on_disconnect() {
let addr = spawn_test_server(RealtimeConfig::default(), TestValidator::new()).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-cleanup"))).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let parsed = parse_server_msg(&msg);
assert_eq!(parsed["type"], "connected");
ws.close(None).await.ok();
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let mut connected = false;
while tokio::time::Instant::now() < deadline {
if let Ok((mut ws2, _)) = connect_async(ws_url(addr, Some("valid-cleanup"))).await {
ws2.close(None).await.ok();
connected = true;
break;
}
tokio::task::yield_now().await;
}
assert!(connected, "reconnect should succeed after disconnect cleanup");
}
async fn spawn_router_test_server(validator: TestValidator) -> SocketAddr {
let config = RealtimeConfig {
max_connections_per_context: 10,
..RealtimeConfig::default()
};
let server = Arc::new(RealtimeServer::with_entities(config, test_entities()));
let state = RealtimeState {
server,
validator: Arc::new(validator),
};
let app = realtime_router(state);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app.into_make_service()).await.unwrap();
});
addr
}
#[tokio::test]
async fn test_realtime_route_mounts_at_path() {
let addr = spawn_router_test_server(TestValidator::new()).await;
let (mut ws, response) = connect_async(ws_url(addr, Some("valid-route-test")))
.await
.expect("WebSocket upgrade should succeed");
assert_eq!(response.status(), 101, "Expected 101 Switching Protocols");
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "connected");
assert!(msg["connection_id"].is_string());
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_realtime_route_without_upgrade_returns_400() {
let addr = spawn_router_test_server(TestValidator::new()).await;
let url = format!("http://127.0.0.1:{}/realtime/v1?token=valid-plain", addr.port());
let response = reqwest::get(&url).await.expect("HTTP request should complete");
assert_eq!(
response.status(),
reqwest::StatusCode::BAD_REQUEST,
"Expected 400 Bad Request for non-WebSocket GET"
);
}
#[test]
fn test_realtime_config_loaded_from_schema() {
let json = serde_json::json!({
"enabled": true,
"entities": ["Post", "Comment"],
"max_connections_per_context": 25
});
let config: RealtimeSchemaConfig = serde_json::from_value(json).unwrap();
assert!(config.enabled, "expected enabled=true");
assert_eq!(config.entities, vec!["Post", "Comment"], "expected entities list");
assert_eq!(config.max_connections_per_context, Some(25));
}
#[test]
fn test_realtime_disabled_in_schema() {
let json = serde_json::json!({ "enabled": false });
let config: RealtimeSchemaConfig = serde_json::from_value(json).unwrap();
assert!(!config.enabled, "expected enabled=false");
assert!(config.entities.is_empty(), "entities should default to empty");
}
#[test]
fn test_realtime_entities_from_schema() {
let json = serde_json::json!({
"enabled": true,
"entities": ["Post", "Comment"]
});
let schema_config: RealtimeSchemaConfig = serde_json::from_value(json).unwrap();
let entities: HashSet<String> = schema_config.entities.into_iter().collect();
let server = Arc::new(RealtimeServer::with_entities(RealtimeConfig::default(), entities));
assert!(server.known_entities.contains("Post"));
assert!(server.known_entities.contains("Comment"));
assert!(!server.known_entities.contains("User"), "User not declared in schema");
assert_eq!(server.known_entities.len(), 2);
}
async fn spawn_combined_server(validator: TestValidator) -> SocketAddr {
let rt_server =
Arc::new(RealtimeServer::with_entities(RealtimeConfig::default(), test_entities()));
let rt_state = RealtimeState {
server: rt_server,
validator: Arc::new(validator),
};
let realtime_app = realtime_router(rt_state);
let schema = Arc::new(CompiledSchema::default());
let sub_manager = Arc::new(SubscriptionManager::new(schema));
let sub_state = crate::routes::subscriptions::SubscriptionState::new(sub_manager);
let graphql_ws_app = Router::new()
.route("/ws", get(crate::routes::subscriptions::subscription_handler))
.with_state(sub_state);
let app = Router::new().merge(realtime_app).merge(graphql_ws_app);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app.into_make_service()).await.unwrap();
});
addr
}
async fn graphql_ws_init(
ws: &mut tokio_tungstenite::WebSocketStream<
tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>,
>,
) -> serde_json::Value {
let init = serde_json::json!({"type": "connection_init"}).to_string();
ws.send(tungstenite::Message::Text(init.into())).await.unwrap();
let msg = ws.next().await.unwrap().unwrap();
let text = match msg {
tungstenite::Message::Text(t) => t.to_string(),
other => panic!("Expected text message, got {other:?}"),
};
serde_json::from_str(&text).unwrap()
}
#[tokio::test]
async fn test_graphql_subscriptions_still_work() {
let addr = spawn_combined_server(TestValidator::new()).await;
let url = format!("ws://127.0.0.1:{}/ws", addr.port());
let mut req = url.into_client_request().unwrap();
req.headers_mut()
.insert("Sec-WebSocket-Protocol", "graphql-transport-ws".parse().unwrap());
let (mut ws, response) = connect_async(req).await.expect("GraphQL WS upgrade should succeed");
assert_eq!(response.status(), 101, "Expected 101 Switching Protocols on /ws");
let ack = graphql_ws_init(&mut ws).await;
assert_eq!(ack["type"], "connection_ack", "Expected connection_ack from graphql-ws handler");
ws.close(None).await.ok();
}
#[tokio::test]
async fn test_realtime_and_graphql_ws_coexist() {
let addr = spawn_combined_server(TestValidator::new()).await;
let (mut rt_ws, rt_resp) = connect_async(ws_url(addr, Some("valid-coexist")))
.await
.expect("Realtime WS upgrade should succeed");
assert_eq!(rt_resp.status(), 101);
let rt_msg = next_msg(&mut rt_ws).await;
assert_eq!(rt_msg["type"], "connected", "Realtime endpoint must send connected message");
let url = format!("ws://127.0.0.1:{}/ws", addr.port());
let mut req = url.into_client_request().unwrap();
req.headers_mut()
.insert("Sec-WebSocket-Protocol", "graphql-transport-ws".parse().unwrap());
let (mut gql_ws, gql_resp) =
connect_async(req).await.expect("GraphQL WS upgrade should succeed");
assert_eq!(gql_resp.status(), 101);
let ack = graphql_ws_init(&mut gql_ws).await;
assert_eq!(ack["type"], "connection_ack");
rt_ws.close(None).await.ok();
gql_ws.close(None).await.ok();
}
#[tokio::test]
async fn test_realtime_route_does_not_conflict() {
let addr = spawn_combined_server(TestValidator::new()).await;
let (mut ws, _) = connect_async(ws_url(addr, Some("valid-no-conflict")))
.await
.expect("Realtime WS upgrade should succeed");
let msg = next_msg(&mut ws).await;
assert_eq!(msg["type"], "connected", "Must be realtime protocol, not graphql-ws");
assert!(
msg["connection_id"].is_string(),
"Realtime protocol always includes connection_id in the connected message"
);
let http_url = format!("http://127.0.0.1:{}/ws", addr.port());
let resp = reqwest::get(&http_url).await.unwrap();
assert_eq!(
resp.status(),
reqwest::StatusCode::BAD_REQUEST,
"/ws without upgrade must return 400, confirming it is a separate route"
);
ws.close(None).await.ok();
}