Update: 将子项目从 submodule 转为完整内容
- 移除 GovAI, nomifun-tauri, 算力盒子 的 submodule 引用 - 添加所有子项目的完整源代码 - 保留原始 .git 为 .git.bak 备份
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use nomifun_api_types::WebSocketMessage;
|
||||
use nomifun_realtime::{BroadcastEventBus, EventBroadcaster};
|
||||
use serde_json::json;
|
||||
|
||||
#[tokio::test]
|
||||
async fn broadcast_to_multiple_subscribers() {
|
||||
let bus = Arc::new(BroadcastEventBus::new(64));
|
||||
let mut rx1 = bus.subscribe();
|
||||
let mut rx2 = bus.subscribe();
|
||||
let mut rx3 = bus.subscribe();
|
||||
|
||||
let event = WebSocketMessage::new("test:broadcast", json!({"key": "value"}));
|
||||
bus.broadcast(event);
|
||||
|
||||
let msg1 = rx1.recv().await.unwrap();
|
||||
let msg2 = rx2.recv().await.unwrap();
|
||||
let msg3 = rx3.recv().await.unwrap();
|
||||
|
||||
assert_eq!(msg1.name, "test:broadcast");
|
||||
assert_eq!(msg2.name, "test:broadcast");
|
||||
assert_eq!(msg3.name, "test:broadcast");
|
||||
assert_eq!(msg1.data, msg2.data);
|
||||
assert_eq!(msg2.data, msg3.data);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn late_subscriber_misses_earlier_events() {
|
||||
let bus = BroadcastEventBus::new(64);
|
||||
|
||||
// Broadcast before any subscriber exists
|
||||
bus.broadcast(WebSocketMessage::new("early", json!({})));
|
||||
|
||||
// Subscribe after the broadcast
|
||||
let mut rx = bus.subscribe();
|
||||
|
||||
// Broadcast a new event
|
||||
bus.broadcast(WebSocketMessage::new("late", json!({})));
|
||||
|
||||
let msg = rx.recv().await.unwrap();
|
||||
assert_eq!(msg.name, "late");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dropped_subscriber_does_not_block_broadcast() {
|
||||
let bus = BroadcastEventBus::new(64);
|
||||
let rx = bus.subscribe();
|
||||
assert_eq!(bus.receiver_count(), 1);
|
||||
|
||||
drop(rx);
|
||||
assert_eq!(bus.receiver_count(), 0);
|
||||
|
||||
// Broadcast should succeed without panic
|
||||
bus.broadcast(WebSocketMessage::new("after-drop", json!({})));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn trait_object_via_arc() {
|
||||
let bus = Arc::new(BroadcastEventBus::new(64));
|
||||
let mut rx = bus.subscribe();
|
||||
|
||||
let broadcaster: Arc<dyn EventBroadcaster> = bus.clone();
|
||||
broadcaster.broadcast(WebSocketMessage::new("via-trait", json!({"n": 42})));
|
||||
|
||||
let msg = rx.recv().await.unwrap();
|
||||
assert_eq!(msg.name, "via-trait");
|
||||
assert_eq!(msg.data["n"], 42);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn high_throughput_broadcast() {
|
||||
let bus = Arc::new(BroadcastEventBus::new(256));
|
||||
let mut rx = bus.subscribe();
|
||||
|
||||
let count = 100;
|
||||
for i in 0..count {
|
||||
bus.broadcast(WebSocketMessage::new(format!("evt-{i}"), json!({"seq": i})));
|
||||
}
|
||||
|
||||
for i in 0..count {
|
||||
let msg = rx.recv().await.unwrap();
|
||||
assert_eq!(msg.name, format!("evt-{i}"));
|
||||
assert_eq!(msg.data["seq"], i);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,395 @@
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::Router;
|
||||
use axum::routing::get;
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use nomifun_api_types::WebSocketMessage;
|
||||
use nomifun_realtime::{
|
||||
ConnectionId, MessageRouter, NoopMessageRouter, WebSocketManager, WsHandlerState, ws_upgrade_handler,
|
||||
};
|
||||
use serde_json::{Value, json};
|
||||
use tokio::net::TcpListener;
|
||||
use tokio_tungstenite::tungstenite;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Test helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Start an axum server with the WebSocket handler and return its address.
|
||||
async fn start_server(state: WsHandlerState) -> SocketAddr {
|
||||
let app = Router::new().route("/ws", get(ws_upgrade_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 default_state() -> (WsHandlerState, Arc<WebSocketManager>) {
|
||||
let manager = Arc::new(WebSocketManager::new());
|
||||
let state = WsHandlerState {
|
||||
manager: manager.clone(),
|
||||
router: Arc::new(NoopMessageRouter),
|
||||
token_validator: Arc::new(|t| t == "valid-token"),
|
||||
token_extractor: Arc::new(|headers| {
|
||||
headers
|
||||
.get("authorization")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.strip_prefix("Bearer "))
|
||||
.map(|s| s.to_owned())
|
||||
}),
|
||||
};
|
||||
(state, manager)
|
||||
}
|
||||
|
||||
/// Connect with an Authorization header.
|
||||
async fn connect_with_token(
|
||||
addr: SocketAddr,
|
||||
token: &str,
|
||||
) -> (
|
||||
futures_util::stream::SplitSink<
|
||||
tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>,
|
||||
tungstenite::Message,
|
||||
>,
|
||||
futures_util::stream::SplitStream<
|
||||
tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>,
|
||||
>,
|
||||
) {
|
||||
let url = format!("ws://{addr}/ws");
|
||||
let request = tungstenite::http::Request::builder()
|
||||
.uri(&url)
|
||||
.header("Host", addr.to_string())
|
||||
.header("Connection", "Upgrade")
|
||||
.header("Upgrade", "websocket")
|
||||
.header("Sec-WebSocket-Version", "13")
|
||||
.header("Sec-WebSocket-Key", tungstenite::handshake::client::generate_key())
|
||||
.header("Authorization", format!("Bearer {token}"))
|
||||
.body(())
|
||||
.unwrap();
|
||||
|
||||
let (ws, _) = tokio_tungstenite::connect_async(request).await.unwrap();
|
||||
ws.split()
|
||||
}
|
||||
|
||||
/// Connect without any auth header.
|
||||
async fn connect_no_token(
|
||||
addr: SocketAddr,
|
||||
) -> tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>> {
|
||||
let url = format!("ws://{addr}/ws");
|
||||
let (ws, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
|
||||
ws
|
||||
}
|
||||
|
||||
/// Read the next text message within a timeout.
|
||||
async fn read_text<S>(stream: &mut S) -> Value
|
||||
where
|
||||
S: StreamExt<Item = Result<tungstenite::Message, tungstenite::Error>> + Unpin,
|
||||
{
|
||||
let timeout = Duration::from_secs(5);
|
||||
tokio::time::timeout(timeout, async {
|
||||
loop {
|
||||
match stream.next().await {
|
||||
Some(Ok(tungstenite::Message::Text(t))) => {
|
||||
return serde_json::from_str::<Value>(&t).unwrap();
|
||||
}
|
||||
Some(Ok(tungstenite::Message::Close(_))) => {
|
||||
panic!("unexpected close frame");
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
panic!("read error: {e}");
|
||||
}
|
||||
None => {
|
||||
panic!("stream ended");
|
||||
}
|
||||
_ => continue, // skip ping/pong/binary
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("read timed out")
|
||||
}
|
||||
|
||||
/// Read until a close frame is received, returning the close code.
|
||||
async fn read_close<S>(stream: &mut S) -> Option<u16>
|
||||
where
|
||||
S: StreamExt<Item = Result<tungstenite::Message, tungstenite::Error>> + Unpin,
|
||||
{
|
||||
let timeout = Duration::from_secs(5);
|
||||
tokio::time::timeout(timeout, async {
|
||||
loop {
|
||||
match stream.next().await {
|
||||
Some(Ok(tungstenite::Message::Close(frame))) => {
|
||||
return frame.map(|f| f.code.into());
|
||||
}
|
||||
Some(Ok(_)) => continue,
|
||||
Some(Err(_)) => return None,
|
||||
None => return None,
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("read_close timed out")
|
||||
}
|
||||
|
||||
fn send_json(text: &str) -> tungstenite::Message {
|
||||
tungstenite::Message::Text(text.into())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn valid_token_connects_successfully() {
|
||||
let (state, manager) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (_tx, _rx) = connect_with_token(addr, "valid-token").await;
|
||||
|
||||
// Allow connection to register
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
assert_eq!(manager.client_count(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn no_token_closes_with_1008() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let mut ws = connect_no_token(addr).await;
|
||||
|
||||
let code = read_close(&mut ws).await;
|
||||
assert_eq!(code, Some(1008));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_token_sends_auth_expired_then_closes() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (_, mut rx) = connect_with_token(addr, "bad-token").await;
|
||||
|
||||
let msg = read_text(&mut rx).await;
|
||||
assert_eq!(msg["name"], "auth-expired");
|
||||
|
||||
let code = read_close(&mut rx).await;
|
||||
assert_eq!(code, Some(1008));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_json_message_returns_error() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, mut rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
tx.send(send_json("not valid json")).await.unwrap();
|
||||
|
||||
let msg = read_text(&mut rx).await;
|
||||
assert_eq!(msg["error"], "Invalid message format");
|
||||
assert!(msg["expected"].is_string());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn missing_fields_returns_error() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, mut rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
tx.send(send_json(r#"{"foo":"bar"}"#)).await.unwrap();
|
||||
|
||||
let msg = read_text(&mut rx).await;
|
||||
assert_eq!(msg["error"], "Invalid message format");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn subscribe_show_open_replies_with_show_open_request() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, mut rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
// Mirrors the @office-ai/platform bridge envelope shape produced by
|
||||
// `invoke('show-open', { properties: ['openFile'] })`.
|
||||
let payload = json!({
|
||||
"name": "subscribe-show-open",
|
||||
"data": {"id": "abc123", "data": {"properties": ["openFile"]}}
|
||||
});
|
||||
tx.send(send_json(&payload.to_string())).await.unwrap();
|
||||
|
||||
let msg = read_text(&mut rx).await;
|
||||
assert_eq!(msg["name"], "show-open-request");
|
||||
assert_eq!(msg["data"]["id"], "abc123");
|
||||
assert_eq!(msg["data"]["isFileMode"], true);
|
||||
assert_eq!(msg["data"]["properties"], json!(["openFile"]));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn subscribe_show_open_directory_mode() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, mut rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
let payload = json!({
|
||||
"name": "subscribe-show-open",
|
||||
"data": {"id": "dir1", "data": {"properties": ["openFile", "openDirectory"]}}
|
||||
});
|
||||
tx.send(send_json(&payload.to_string())).await.unwrap();
|
||||
|
||||
let msg = read_text(&mut rx).await;
|
||||
assert_eq!(msg["name"], "show-open-request");
|
||||
assert_eq!(msg["data"]["id"], "dir1");
|
||||
assert_eq!(msg["data"]["isFileMode"], false);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn broadcast_reaches_all_connected_clients() {
|
||||
let (state, manager) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (_, mut rx1) = connect_with_token(addr, "valid-token").await;
|
||||
let (_, mut rx2) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
assert_eq!(manager.client_count(), 2);
|
||||
|
||||
let event = WebSocketMessage::new("test-broadcast", json!({"seq": 1}));
|
||||
manager.broadcast_all(event);
|
||||
|
||||
let msg1 = read_text(&mut rx1).await;
|
||||
let msg2 = read_text(&mut rx2).await;
|
||||
|
||||
assert_eq!(msg1["name"], "test-broadcast");
|
||||
assert_eq!(msg2["name"], "test-broadcast");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unicast_reaches_only_target() {
|
||||
let (state, manager) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (_, mut rx1) = connect_with_token(addr, "valid-token").await;
|
||||
let (_, mut rx2) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
assert_eq!(manager.client_count(), 2);
|
||||
|
||||
// IDs are sequential starting from 1
|
||||
let first_conn_id = ConnectionId(1);
|
||||
|
||||
let msg = WebSocketMessage::new("unicast-test", json!({"target": true}));
|
||||
manager.send_to(first_conn_id, msg);
|
||||
|
||||
let received = read_text(&mut rx1).await;
|
||||
assert_eq!(received["name"], "unicast-test");
|
||||
|
||||
// rx2 should not have received anything — check with short timeout
|
||||
let timeout_result = tokio::time::timeout(Duration::from_millis(200), rx2.next()).await;
|
||||
assert!(timeout_result.is_err(), "rx2 should not receive the unicast");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn client_disconnect_removes_from_manager() {
|
||||
let (state, manager) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, _rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
assert_eq!(manager.client_count(), 1);
|
||||
|
||||
// Send close frame
|
||||
tx.send(tungstenite::Message::Close(None)).await.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
|
||||
assert_eq!(manager.client_count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pong_message_does_not_generate_response() {
|
||||
let (state, _) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let (mut tx, mut rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
let pong = json!({"name": "pong", "data": {}});
|
||||
tx.send(send_json(&pong.to_string())).await.unwrap();
|
||||
|
||||
// pong should not generate any response
|
||||
let timeout_result = tokio::time::timeout(Duration::from_millis(200), rx.next()).await;
|
||||
assert!(timeout_result.is_err(), "pong should not generate a response");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unknown_message_routed_to_message_router() {
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
struct TrackingRouter {
|
||||
called: AtomicBool,
|
||||
}
|
||||
impl MessageRouter for TrackingRouter {
|
||||
fn route(&self, _conn_id: ConnectionId, _name: &str, _data: Value) {
|
||||
self.called.store(true, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
let manager = Arc::new(WebSocketManager::new());
|
||||
let router = Arc::new(TrackingRouter {
|
||||
called: AtomicBool::new(false),
|
||||
});
|
||||
let state = WsHandlerState {
|
||||
manager: manager.clone(),
|
||||
router: router.clone(),
|
||||
token_validator: Arc::new(|t| t == "valid-token"),
|
||||
token_extractor: Arc::new(|headers| {
|
||||
headers
|
||||
.get("authorization")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.strip_prefix("Bearer "))
|
||||
.map(|s| s.to_owned())
|
||||
}),
|
||||
};
|
||||
|
||||
let addr = start_server(state).await;
|
||||
let (mut tx, _rx) = connect_with_token(addr, "valid-token").await;
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
let msg = json!({"name": "custom.business-event", "data": {"key": "val"}});
|
||||
tx.send(send_json(&msg.to_string())).await.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
|
||||
assert!(router.called.load(Ordering::Relaxed));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn multiple_concurrent_connections() {
|
||||
let (state, manager) = default_state();
|
||||
let addr = start_server(state).await;
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for _ in 0..10 {
|
||||
handles.push(tokio::spawn(
|
||||
async move { connect_with_token(addr, "valid-token").await },
|
||||
));
|
||||
}
|
||||
|
||||
let mut connections = Vec::new();
|
||||
for h in handles {
|
||||
connections.push(h.await.unwrap());
|
||||
}
|
||||
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
assert_eq!(manager.client_count(), 10);
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use nomifun_api_types::WebSocketMessage;
|
||||
use nomifun_realtime::{
|
||||
ConnectionId, PER_CONNECTION_BUFFER, TokenValidator, WebSocketCloseCode, WebSocketManager, WsOutbound,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
fn always_valid() -> TokenValidator {
|
||||
Arc::new(|_| true)
|
||||
}
|
||||
|
||||
fn new_client_tx() -> (mpsc::Sender<WsOutbound>, mpsc::Receiver<WsOutbound>) {
|
||||
mpsc::channel(PER_CONNECTION_BUFFER)
|
||||
}
|
||||
|
||||
// --- Connection lifecycle ---
|
||||
|
||||
#[test]
|
||||
fn register_and_remove_multiple_clients() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let mut ids = Vec::new();
|
||||
|
||||
for i in 0..10 {
|
||||
let (tx, _rx) = new_client_tx();
|
||||
let id = mgr.add_client(format!("token-{i}"), tx);
|
||||
ids.push(id);
|
||||
}
|
||||
|
||||
assert_eq!(mgr.client_count(), 10);
|
||||
|
||||
// Remove every other client
|
||||
for id in ids.iter().step_by(2) {
|
||||
mgr.remove_client(*id);
|
||||
}
|
||||
assert_eq!(mgr.client_count(), 5);
|
||||
|
||||
// Remove remaining
|
||||
for id in ids.iter().skip(1).step_by(2) {
|
||||
mgr.remove_client(*id);
|
||||
}
|
||||
assert_eq!(mgr.client_count(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_ids_are_unique_and_monotonic() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let mut ids = Vec::new();
|
||||
|
||||
for _ in 0..100 {
|
||||
let (tx, _rx) = new_client_tx();
|
||||
ids.push(mgr.add_client("tok".into(), tx));
|
||||
}
|
||||
|
||||
// Check uniqueness
|
||||
let mut sorted = ids.clone();
|
||||
sorted.sort_by_key(|id| id.0);
|
||||
sorted.dedup();
|
||||
assert_eq!(sorted.len(), 100);
|
||||
|
||||
// Check monotonic
|
||||
for window in ids.windows(2) {
|
||||
assert!(window[0].0 < window[1].0);
|
||||
}
|
||||
}
|
||||
|
||||
// --- Broadcast ---
|
||||
|
||||
#[test]
|
||||
fn broadcast_all_delivers_identical_content_to_every_client() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let mut receivers = Vec::new();
|
||||
|
||||
for i in 0..5 {
|
||||
let (tx, rx) = new_client_tx();
|
||||
mgr.add_client(format!("token-{i}"), tx);
|
||||
receivers.push(rx);
|
||||
}
|
||||
|
||||
let event = WebSocketMessage::new("notification", json!({"level": "info", "text": "hello"}));
|
||||
mgr.broadcast_all(event);
|
||||
|
||||
let mut texts = Vec::new();
|
||||
for rx in &mut receivers {
|
||||
match rx.try_recv().unwrap() {
|
||||
WsOutbound::Text(t) => texts.push(t),
|
||||
other => panic!("expected Text, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
// All received identical content
|
||||
assert!(texts.windows(2).all(|w| w[0] == w[1]));
|
||||
assert!(texts[0].contains("notification"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn broadcast_cleans_up_disconnected_clients_transparently() {
|
||||
let mgr = WebSocketManager::new();
|
||||
|
||||
// 3 live clients
|
||||
let (tx1, _rx1) = new_client_tx();
|
||||
let (tx2, _rx2) = new_client_tx();
|
||||
let (tx3, _rx3) = new_client_tx();
|
||||
mgr.add_client("a".into(), tx1);
|
||||
mgr.add_client("b".into(), tx2);
|
||||
mgr.add_client("c".into(), tx3);
|
||||
|
||||
// 2 dead clients (receivers dropped)
|
||||
let (tx4, rx4) = new_client_tx();
|
||||
let (tx5, rx5) = new_client_tx();
|
||||
mgr.add_client("dead-1".into(), tx4);
|
||||
mgr.add_client("dead-2".into(), tx5);
|
||||
drop(rx4);
|
||||
drop(rx5);
|
||||
|
||||
assert_eq!(mgr.client_count(), 5);
|
||||
|
||||
mgr.broadcast_all(WebSocketMessage::new("check", json!(null)));
|
||||
|
||||
// Dead clients should be removed
|
||||
assert_eq!(mgr.client_count(), 3);
|
||||
}
|
||||
|
||||
// --- Unicast ---
|
||||
|
||||
#[test]
|
||||
fn send_to_reaches_only_target_connection() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let mut pairs: Vec<(ConnectionId, mpsc::Receiver<WsOutbound>)> = Vec::new();
|
||||
|
||||
for i in 0..5 {
|
||||
let (tx, rx) = new_client_tx();
|
||||
let id = mgr.add_client(format!("token-{i}"), tx);
|
||||
pairs.push((id, rx));
|
||||
}
|
||||
|
||||
let target_id = pairs[2].0;
|
||||
mgr.send_to(target_id, WebSocketMessage::new("private", json!({"secret": true})));
|
||||
|
||||
for (id, rx) in &mut pairs {
|
||||
if *id == target_id {
|
||||
let msg = rx.try_recv().unwrap();
|
||||
match msg {
|
||||
WsOutbound::Text(t) => assert!(t.contains("private")),
|
||||
other => panic!("expected Text, got {other:?}"),
|
||||
}
|
||||
} else {
|
||||
assert!(rx.try_recv().is_err(), "non-target {id} should not receive message");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Heartbeat integration ---
|
||||
|
||||
#[tokio::test]
|
||||
async fn heartbeat_sends_ping_and_keeps_healthy_connections() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let (tx, mut rx) = new_client_tx();
|
||||
mgr.add_client("valid-token".into(), tx);
|
||||
|
||||
let handle = mgr.start_heartbeat(always_valid());
|
||||
|
||||
// Wait for first heartbeat tick (interval is 30s, but first tick fires immediately)
|
||||
let msg = tokio::time::timeout(Duration::from_secs(2), rx.recv())
|
||||
.await
|
||||
.expect("timeout waiting for ping")
|
||||
.expect("channel closed");
|
||||
|
||||
match msg {
|
||||
WsOutbound::Text(text) => {
|
||||
let parsed: serde_json::Value = serde_json::from_str(&text).unwrap();
|
||||
assert_eq!(parsed["name"], "ping");
|
||||
assert!(parsed["data"]["timestamp"].is_u64());
|
||||
}
|
||||
other => panic!("expected ping Text, got {other:?}"),
|
||||
}
|
||||
|
||||
// Connection should still be alive
|
||||
assert_eq!(mgr.client_count(), 1);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn heartbeat_closes_expired_token_with_auth_expired_event() {
|
||||
let mgr = WebSocketManager::new();
|
||||
let (tx, mut rx) = new_client_tx();
|
||||
mgr.add_client("bad-token".into(), tx);
|
||||
|
||||
let expired_validator: TokenValidator = Arc::new(|_| false);
|
||||
let handle = mgr.start_heartbeat(expired_validator);
|
||||
|
||||
// Expect auth-expired event
|
||||
let msg1 = tokio::time::timeout(Duration::from_secs(2), rx.recv())
|
||||
.await
|
||||
.expect("timeout")
|
||||
.expect("closed");
|
||||
|
||||
match msg1 {
|
||||
WsOutbound::Text(text) => {
|
||||
let parsed: serde_json::Value = serde_json::from_str(&text).unwrap();
|
||||
assert_eq!(parsed["name"], "auth-expired");
|
||||
assert!(parsed["data"]["message"].is_string());
|
||||
}
|
||||
other => panic!("expected auth-expired, got {other:?}"),
|
||||
}
|
||||
|
||||
// Expect close frame
|
||||
let msg2 = tokio::time::timeout(Duration::from_secs(1), rx.recv())
|
||||
.await
|
||||
.expect("timeout")
|
||||
.expect("closed");
|
||||
|
||||
assert_eq!(
|
||||
msg2,
|
||||
WsOutbound::Close(WebSocketCloseCode::PolicyViolation, "token expired".into())
|
||||
);
|
||||
|
||||
// Connection should be removed
|
||||
assert_eq!(mgr.client_count(), 0);
|
||||
|
||||
handle.abort();
|
||||
}
|
||||
|
||||
// --- Concurrent access ---
|
||||
|
||||
#[test]
|
||||
fn concurrent_add_remove_does_not_panic() {
|
||||
let mgr = Arc::new(WebSocketManager::new());
|
||||
let mut handles = Vec::new();
|
||||
|
||||
// Spawn threads that add clients
|
||||
for i in 0..10 {
|
||||
let mgr = Arc::clone(&mgr);
|
||||
handles.push(std::thread::spawn(move || {
|
||||
let (tx, _rx) = new_client_tx();
|
||||
mgr.add_client(format!("thread-{i}"), tx)
|
||||
}));
|
||||
}
|
||||
|
||||
let ids: Vec<ConnectionId> = handles.into_iter().map(|h| h.join().unwrap()).collect();
|
||||
|
||||
assert_eq!(mgr.client_count(), 10);
|
||||
|
||||
// All IDs should be unique
|
||||
let mut unique = ids.clone();
|
||||
unique.sort_by_key(|id| id.0);
|
||||
unique.dedup();
|
||||
assert_eq!(unique.len(), 10);
|
||||
|
||||
// Remove all concurrently
|
||||
let mut handles = Vec::new();
|
||||
for id in ids {
|
||||
let mgr = Arc::clone(&mgr);
|
||||
handles.push(std::thread::spawn(move || {
|
||||
mgr.remove_client(id);
|
||||
}));
|
||||
}
|
||||
|
||||
for h in handles {
|
||||
h.join().unwrap();
|
||||
}
|
||||
|
||||
assert_eq!(mgr.client_count(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn concurrent_broadcast_does_not_panic() {
|
||||
let mgr = Arc::new(WebSocketManager::new());
|
||||
let mut _receivers = Vec::new();
|
||||
|
||||
for i in 0..5 {
|
||||
let (tx, rx) = new_client_tx();
|
||||
mgr.add_client(format!("tok-{i}"), tx);
|
||||
_receivers.push(rx);
|
||||
}
|
||||
|
||||
let mut handles = Vec::new();
|
||||
for i in 0..10 {
|
||||
let mgr = Arc::clone(&mgr);
|
||||
handles.push(std::thread::spawn(move || {
|
||||
mgr.broadcast_all(WebSocketMessage::new(format!("event-{i}"), json!(null)));
|
||||
}));
|
||||
}
|
||||
|
||||
for h in handles {
|
||||
h.join().unwrap();
|
||||
}
|
||||
|
||||
// All clients should still be connected
|
||||
assert_eq!(mgr.client_count(), 5);
|
||||
}
|
||||
Reference in New Issue
Block a user