Skip to content

Commit e0a5cfb

Browse files
committed
chore: Use real origin format in tests and assert X-Cll-Origin header is sent
1 parent a5d77b5 commit e0a5cfb

3 files changed

Lines changed: 100 additions & 28 deletions

File tree

rust/crates/sdk/src/stream.rs

Lines changed: 6 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -400,36 +400,20 @@ mod tests {
400400

401401
#[test]
402402
fn test_parse_origins_from_header_with_braces() {
403-
let result = parse_origins_from_header(
404-
"{wss://ws1.dataengine.chain.link,wss://ws2.dataengine.chain.link}",
405-
);
406-
assert_eq!(
407-
result,
408-
vec![
409-
"wss://ws1.dataengine.chain.link".to_string(),
410-
"wss://ws2.dataengine.chain.link".to_string(),
411-
]
412-
);
403+
let result = parse_origins_from_header("{001,002}");
404+
assert_eq!(result, vec!["001".to_string(), "002".to_string()]);
413405
}
414406

415407
#[test]
416408
fn test_parse_origins_from_header_without_braces() {
417-
let result = parse_origins_from_header(
418-
"wss://ws1.dataengine.chain.link,wss://ws2.dataengine.chain.link",
419-
);
420-
assert_eq!(
421-
result,
422-
vec![
423-
"wss://ws1.dataengine.chain.link".to_string(),
424-
"wss://ws2.dataengine.chain.link".to_string(),
425-
]
426-
);
409+
let result = parse_origins_from_header("001,002");
410+
assert_eq!(result, vec!["001".to_string(), "002".to_string()]);
427411
}
428412

429413
#[test]
430414
fn test_parse_origins_from_header_single_origin() {
431-
let result = parse_origins_from_header("{wss://ws1.dataengine.chain.link}");
432-
assert_eq!(result, vec!["wss://ws1.dataengine.chain.link".to_string()]);
415+
let result = parse_origins_from_header("{001}");
416+
assert_eq!(result, vec!["001".to_string()]);
433417
}
434418

435419
#[test]

rust/crates/sdk/tests/stream_integration_tests.rs

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -263,6 +263,59 @@ async fn test_stream_ha_reconnect_merge() {
263263
assert_eq!(stats.deduplicated, expected_deduplicated);
264264
}
265265

266+
#[tokio::test]
267+
async fn test_stream_ha_x_cll_origin_header() {
268+
let mock_server = MockWebSocketServer::new("127.0.0.1:0").await;
269+
let ws_url = format!("ws://{}", mock_server.address());
270+
271+
// Use distinct opaque origin IDs matching the real protocol format (e.g. {001,002}).
272+
mock_server
273+
.set_ha_origins(vec!["001".to_string(), "002".to_string()])
274+
.await;
275+
276+
let config = Config::new(
277+
"mock_key".to_string(),
278+
"mock_secret".to_string(),
279+
"mock_rest_url".to_string(),
280+
ws_url,
281+
)
282+
.with_ws_ha(WebSocketHighAvailability::Enabled)
283+
.with_ws_max_reconnect(MAX_RECONNECT_ATTEMPTS)
284+
.build()
285+
.expect("Failed to build config");
286+
287+
let mut stream = Stream::new(&config, vec![])
288+
.await
289+
.expect("Failed to create stream");
290+
291+
stream.listen().await.expect("Failed to start listening");
292+
293+
// Allow time for both WebSocket connections to be established.
294+
sleep(Duration::from_millis(500)).await;
295+
296+
let received = mock_server.get_received_cll_origins().await;
297+
298+
// Assert 1: X-Cll-Origin header is present on every HA WebSocket connection.
299+
assert_eq!(received.len(), 2, "Expected 2 WebSocket connections in HA mode");
300+
for origin in &received {
301+
assert!(
302+
origin.is_some(),
303+
"X-Cll-Origin header was missing on a WebSocket connection"
304+
);
305+
}
306+
307+
// Assert 2: The server sees distinct origin values matching the configured origins.
308+
let mut actual: Vec<String> = received.into_iter().flatten().collect();
309+
actual.sort();
310+
assert_eq!(
311+
actual,
312+
vec!["001".to_string(), "002".to_string()],
313+
"X-Cll-Origin header values did not match the configured origins"
314+
);
315+
316+
stream.close().await.expect("Failed to close stream");
317+
}
318+
266319
#[tokio::test]
267320
#[ignore] // Ignored because it takes a while to complete. To run it, use this command: cargo test -- --ignored
268321
async fn test_stream_ha_max_reconnection_attempts() {

rust/crates/sdk/tests/utils/mock_websocket_server.rs

Lines changed: 41 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,12 @@ use std::sync::Arc;
33
use tokio::{
44
io::{AsyncReadExt, AsyncWriteExt},
55
net::{TcpListener, TcpStream},
6-
sync::{mpsc, Mutex, Notify},
6+
sync::{mpsc, oneshot, Mutex, Notify},
7+
};
8+
use tokio_tungstenite::{
9+
accept_hdr_async,
10+
tungstenite::{handshake::server::Request as WsRequest, protocol::Message},
711
};
8-
use tokio_tungstenite::{accept_async, tungstenite::protocol::Message};
912

1013
enum ServerCommand {
1114
Send(Vec<u8>),
@@ -20,6 +23,9 @@ pub struct MockWebSocketServer {
2023
/// Origins returned in X-Cll-Available-Origins HEAD response.
2124
/// When None, defaults to two copies of the server's own ws:// address.
2225
ha_origins: Arc<Mutex<Option<Vec<String>>>>,
26+
/// X-Cll-Origin header values captured from incoming WebSocket upgrade requests.
27+
/// Some(value) if the header was present, None if absent.
28+
received_cll_origins: Arc<Mutex<Vec<Option<String>>>>,
2329
}
2430

2531
impl MockWebSocketServer {
@@ -35,10 +41,13 @@ impl MockWebSocketServer {
3541
let clients = Arc::new(Mutex::new(Vec::new()));
3642
let shutdown_notify = Arc::new(Notify::new());
3743
let ha_origins: Arc<Mutex<Option<Vec<String>>>> = Arc::new(Mutex::new(None));
44+
let received_cll_origins: Arc<Mutex<Vec<Option<String>>>> =
45+
Arc::new(Mutex::new(Vec::new()));
3846

3947
let clients_accept = clients.clone();
4048
let shutdown_accept = shutdown_notify.clone();
4149
let ha_origins_accept = ha_origins.clone();
50+
let received_accept = received_cll_origins.clone();
4251
let server_address = address.clone();
4352

4453
tokio::spawn(async move {
@@ -55,7 +64,8 @@ impl MockWebSocketServer {
5564
])
5665
};
5766
let clients_clone = clients_accept.clone();
58-
tokio::spawn(handle_connection(stream, origins, clients_clone));
67+
let received_clone = received_accept.clone();
68+
tokio::spawn(handle_connection(stream, origins, clients_clone, received_clone));
5969
}
6070
Err(e) => {
6171
println!("Error accepting connection: {:?}", e);
@@ -95,6 +105,7 @@ impl MockWebSocketServer {
95105
command_sender,
96106
shutdown_notify,
97107
ha_origins,
108+
received_cll_origins,
98109
}
99110
}
100111

@@ -122,15 +133,22 @@ impl MockWebSocketServer {
122133
pub async fn set_ha_origins(&self, origins: Vec<String>) {
123134
*self.ha_origins.lock().await = Some(origins);
124135
}
136+
137+
/// Returns the X-Cll-Origin header values captured from all WebSocket upgrade requests.
138+
/// Some(value) means the header was present; None means it was absent.
139+
pub async fn get_received_cll_origins(&self) -> Vec<Option<String>> {
140+
self.received_cll_origins.lock().await.clone()
141+
}
125142
}
126143

127144
async fn handle_connection(
128145
mut stream: TcpStream,
129146
ha_origins: Vec<String>,
130147
clients: Arc<Mutex<Vec<mpsc::Sender<Message>>>>,
148+
received_cll_origins: Arc<Mutex<Vec<Option<String>>>>,
131149
) {
132150
// Peek at first 4 bytes to distinguish HTTP HEAD from WebSocket upgrade.
133-
// peek() does not consume data, so the full request remains readable by accept_async.
151+
// peek() does not consume data, so the full request remains readable by accept_hdr_async.
134152
let mut peek_buf = [0u8; 4];
135153
let n = match stream.peek(&mut peek_buf).await {
136154
Ok(n) => n,
@@ -152,14 +170,31 @@ async fn handle_connection(
152170
);
153171
let _ = stream.write_all(response.as_bytes()).await;
154172
} else {
155-
// WebSocket upgrade
156-
let ws_stream = match accept_async(stream).await {
173+
// WebSocket upgrade — capture the X-Cll-Origin header from the upgrade request.
174+
let (origin_tx, mut origin_rx) = oneshot::channel::<Option<String>>();
175+
176+
let ws_stream = match accept_hdr_async(stream, move |req: &WsRequest, resp| {
177+
let origin = req
178+
.headers()
179+
.get("x-cll-origin")
180+
.and_then(|v| v.to_str().ok())
181+
.map(|s| s.to_string());
182+
let _ = origin_tx.send(origin);
183+
Ok(resp)
184+
})
185+
.await
186+
{
157187
Ok(s) => s,
158188
Err(e) => {
159189
println!("WebSocket accept error: {:?}", e);
160190
return;
161191
}
162192
};
193+
194+
// origin_tx.send() has already run by the time accept_hdr_async resolves.
195+
let origin = origin_rx.try_recv().unwrap_or(None);
196+
received_cll_origins.lock().await.push(origin);
197+
163198
println!(
164199
"Client connected: {}",
165200
ws_stream.get_ref().peer_addr().unwrap()

0 commit comments

Comments
 (0)