@@ -3,9 +3,12 @@ use std::sync::Arc;
33use 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
1013enum 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
2531impl 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
127144async 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