1use super::tcp::stream::validate_stream_record;
9use super::{
10 ConnectionState, DiscoveredPeer, PacketBuffer, PacketTx, ReceivedPacket, Transport,
11 TransportAddr, TransportError, TransportId, TransportState, TransportType,
12};
13use crate::Identity;
14use crate::config::WebSocketConfig;
15use crate::dataplane::validate_direct_fsp_transport_fragment;
16use crate::discovery::local_udp::LocalKeyHint;
17use futures::{SinkExt, StreamExt};
18use rand::RngExt;
19use secp256k1::XOnlyPublicKey;
20use serde::Serialize;
21use std::collections::{HashMap, VecDeque};
22use std::net::SocketAddr;
23use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
24use std::sync::{Arc, Mutex as StdMutex};
25use std::time::{Duration, SystemTime, UNIX_EPOCH};
26use tokio::net::{TcpListener, TcpStream};
27use tokio::sync::{Mutex, Semaphore, mpsc, watch};
28use tokio::task::JoinHandle;
29use tokio_tungstenite::tungstenite::handshake::server::{ErrorResponse, Request, Response};
30use tokio_tungstenite::tungstenite::protocol::WebSocketConfig as TungsteniteConfig;
31use tokio_tungstenite::tungstenite::{Bytes, Message};
32use tokio_tungstenite::{WebSocketStream, accept_hdr_async_with_config, connect_async_with_config};
33use tracing::{debug, info};
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36enum Direction {
37 Inbound,
38 Outbound,
39}
40
41struct Connection {
42 generation: u64,
43 tx: mpsc::Sender<Vec<u8>>,
44}
45
46type ConnectionPool = Arc<Mutex<HashMap<TransportAddr, Connection>>>;
47type ConnectionStates = Arc<StdMutex<HashMap<TransportAddr, ConnectionState>>>;
48type DiscoveryQueue = Arc<StdMutex<VecDeque<DiscoveredPeer>>>;
49
50#[derive(Debug, Default)]
51struct WebSocketStats {
52 connections_opened: AtomicU64,
53 connections_closed: AtomicU64,
54 connections_rejected: AtomicU64,
55 reconnect_attempts: AtomicU64,
56 frames_sent: AtomicU64,
57 frames_received: AtomicU64,
58 bytes_sent: AtomicU64,
59 bytes_received: AtomicU64,
60 invalid_frames: AtomicU64,
61 send_queue_full: AtomicU64,
62}
63
64#[derive(Debug, Serialize)]
65pub struct WebSocketStatsSnapshot {
66 pub connections_opened: u64,
67 pub connections_closed: u64,
68 pub connections_rejected: u64,
69 pub reconnect_attempts: u64,
70 pub frames_sent: u64,
71 pub frames_received: u64,
72 pub bytes_sent: u64,
73 pub bytes_received: u64,
74 pub invalid_frames: u64,
75 pub send_queue_full: u64,
76}
77
78impl WebSocketStats {
79 fn snapshot(&self) -> WebSocketStatsSnapshot {
80 let load = |value: &AtomicU64| value.load(Ordering::Relaxed);
81 WebSocketStatsSnapshot {
82 connections_opened: load(&self.connections_opened),
83 connections_closed: load(&self.connections_closed),
84 connections_rejected: load(&self.connections_rejected),
85 reconnect_attempts: load(&self.reconnect_attempts),
86 frames_sent: load(&self.frames_sent),
87 frames_received: load(&self.frames_received),
88 bytes_sent: load(&self.bytes_sent),
89 bytes_received: load(&self.bytes_received),
90 invalid_frames: load(&self.invalid_frames),
91 send_queue_full: load(&self.send_queue_full),
92 }
93 }
94}
95
96#[derive(Clone)]
97struct Runtime {
98 transport_id: TransportId,
99 config: WebSocketConfig,
100 local_pubkey: [u8; 32],
101 packet_tx: PacketTx,
102 pool: ConnectionPool,
103 states: ConnectionStates,
104 discoveries: DiscoveryQueue,
105 running: Arc<AtomicBool>,
106 total_slots: Arc<Semaphore>,
107 inbound_slots: Arc<Semaphore>,
108 generation: Arc<AtomicU64>,
109 network_rebind_generation: watch::Sender<u64>,
111 stats: Arc<WebSocketStats>,
112}
113
114impl Runtime {
115 fn websocket_config(&self) -> TungsteniteConfig {
116 let mut config = TungsteniteConfig::default();
117 config.max_message_size = Some(self.config.max_frame_bytes());
118 config.max_frame_size = Some(self.config.max_frame_bytes());
119 config.max_write_buffer_size = self.config.max_frame_bytes().saturating_mul(2);
120 config.write_buffer_size = 0;
121 config
122 }
123
124 fn next_generation(&self) -> u64 {
125 self.generation.fetch_add(1, Ordering::Relaxed)
126 }
127
128 fn set_state(&self, addr: &TransportAddr, state: ConnectionState) {
129 self.states
130 .lock()
131 .unwrap_or_else(|error| error.into_inner())
132 .insert(addr.clone(), state);
133 }
134
135 fn clear_state_if(&self, addr: &TransportAddr, generation: u64) {
136 let connected_generation = self
137 .pool
138 .try_lock()
139 .ok()
140 .and_then(|pool| pool.get(addr).map(|connection| connection.generation));
141 if connected_generation != Some(generation) {
142 self.states
143 .lock()
144 .unwrap_or_else(|error| error.into_inner())
145 .remove(addr);
146 }
147 }
148}
149
150pub struct WebSocketTransport {
152 transport_id: TransportId,
153 name: Option<String>,
154 config: WebSocketConfig,
155 state: TransportState,
156 local_addr: Option<SocketAddr>,
157 runtime: Runtime,
158 tasks: Arc<StdMutex<Vec<JoinHandle<()>>>>,
159}
160
161impl WebSocketTransport {
162 pub fn new(
163 transport_id: TransportId,
164 name: Option<String>,
165 config: WebSocketConfig,
166 packet_tx: PacketTx,
167 identity: &Identity,
168 ) -> Self {
169 let max_connections = config.max_connections();
170 let max_inbound = config.max_inbound_connections();
171 let (network_rebind_generation, _) = watch::channel(0);
172 let runtime = Runtime {
173 transport_id,
174 config: config.clone(),
175 local_pubkey: identity.pubkey().serialize(),
176 packet_tx,
177 pool: Arc::new(Mutex::new(HashMap::new())),
178 states: Arc::new(StdMutex::new(HashMap::new())),
179 discoveries: Arc::new(StdMutex::new(VecDeque::new())),
180 running: Arc::new(AtomicBool::new(false)),
181 total_slots: Arc::new(Semaphore::new(max_connections)),
182 inbound_slots: Arc::new(Semaphore::new(max_inbound)),
183 generation: Arc::new(AtomicU64::new(1)),
184 network_rebind_generation,
185 stats: Arc::new(WebSocketStats::default()),
186 };
187 Self {
188 transport_id,
189 name,
190 config,
191 state: TransportState::Configured,
192 local_addr: None,
193 runtime,
194 tasks: Arc::new(StdMutex::new(Vec::new())),
195 }
196 }
197
198 pub fn name(&self) -> Option<&str> {
199 self.name.as_deref()
200 }
201
202 pub fn local_addr(&self) -> Option<SocketAddr> {
203 self.local_addr
204 }
205
206 pub fn public_url(&self) -> Option<&str> {
207 self.config.public_url.as_deref()
208 }
209
210 pub(crate) fn is_configured_seed_addr(&self, addr: &TransportAddr) -> bool {
211 addr.as_str().is_some_and(|candidate| {
212 self.config
213 .seed_urls
214 .iter()
215 .any(|seed_url| seed_url == candidate)
216 })
217 }
218
219 pub(crate) fn is_configured_adjacency(
220 &self,
221 addr: &TransportAddr,
222 handshake_is_initiator: bool,
223 ) -> bool {
224 if self.is_configured_seed_addr(addr) {
225 true
229 } else if !handshake_is_initiator {
230 self.config.bind_addr.is_some()
234 } else {
235 false
236 }
237 }
238
239 pub fn stats(&self) -> WebSocketStatsSnapshot {
240 self.runtime.stats.snapshot()
241 }
242
243 fn push_task(&self, task: JoinHandle<()>) {
244 self.tasks
245 .lock()
246 .unwrap_or_else(|error| error.into_inner())
247 .push(task);
248 }
249
250 pub async fn start_async(&mut self) -> Result<(), TransportError> {
251 if !self.state.can_start() {
252 return Err(TransportError::AlreadyStarted);
253 }
254 self.config
255 .validate()
256 .map_err(TransportError::StartFailed)?;
257 self.state = TransportState::Starting;
258 self.runtime.running.store(true, Ordering::Release);
259
260 if let Some(bind_addr) = self.config.bind_addr.as_deref() {
261 let bind_addr = bind_addr
262 .parse::<SocketAddr>()
263 .map_err(|error| TransportError::StartFailed(error.to_string()))?;
264 let listener = TcpListener::bind(bind_addr)
265 .await
266 .map_err(|error| TransportError::bind_failed(bind_addr, error))?;
267 self.local_addr = Some(
268 listener
269 .local_addr()
270 .map_err(|error| TransportError::StartFailed(error.to_string()))?,
271 );
272 self.push_task(tokio::spawn(run_accept_loop(
273 self.runtime.clone(),
274 listener,
275 )));
276 }
277
278 for seed_url in self.config.seed_urls.clone() {
279 let addr = TransportAddr::from_string(&seed_url);
280 self.runtime.set_state(&addr, ConnectionState::Connecting);
281 self.push_task(tokio::spawn(run_seed_dialer(self.runtime.clone(), addr)));
282 }
283
284 self.state = TransportState::Up;
285 info!(
286 transport_id = %self.transport_id,
287 local_addr = ?self.local_addr,
288 seeds = self.config.seed_urls.len(),
289 "WebSocket transport started"
290 );
291 Ok(())
292 }
293
294 pub async fn stop_async(&mut self) -> Result<(), TransportError> {
295 if !self.state.is_operational() {
296 return Err(TransportError::NotStarted);
297 }
298 self.runtime.running.store(false, Ordering::Release);
299 let tasks = {
300 let mut tasks = self.tasks.lock().unwrap_or_else(|error| error.into_inner());
301 std::mem::take(&mut *tasks)
302 };
303 for task in &tasks {
304 task.abort();
305 }
306 for task in tasks {
307 let _ = task.await;
308 }
309 self.runtime.pool.lock().await.clear();
310 self.runtime
311 .states
312 .lock()
313 .unwrap_or_else(|error| error.into_inner())
314 .clear();
315 self.local_addr = None;
316 self.state = TransportState::Down;
317 Ok(())
318 }
319
320 pub async fn send_async(
321 &self,
322 addr: &TransportAddr,
323 data: &[u8],
324 ) -> Result<usize, TransportError> {
325 if !self.state.is_operational() {
326 return Err(TransportError::NotStarted);
327 }
328 if data.len() > self.config.max_frame_bytes() {
329 return Err(TransportError::MtuExceeded {
330 packet_size: data.len(),
331 mtu: self.config.max_frame_bytes().min(u16::MAX as usize) as u16,
332 });
333 }
334 validate_websocket_record(data).map_err(TransportError::SendFailed)?;
335 let tx = self
336 .runtime
337 .pool
338 .lock()
339 .await
340 .get(addr)
341 .map(|connection| connection.tx.clone())
342 .ok_or(TransportError::NotStarted)?;
343 tx.try_send(data.to_vec()).map_err(|error| match error {
344 mpsc::error::TrySendError::Full(_) => {
345 self.runtime
346 .stats
347 .send_queue_full
348 .fetch_add(1, Ordering::Relaxed);
349 TransportError::SendFailed("WebSocket send queue full".into())
350 }
351 mpsc::error::TrySendError::Closed(_) => {
352 TransportError::SendFailed("WebSocket connection closed".into())
353 }
354 })?;
355 Ok(data.len())
356 }
357
358 pub async fn connect_async(&self, addr: &TransportAddr) -> Result<(), TransportError> {
359 if !self.state.is_operational() {
360 return Err(TransportError::NotStarted);
361 }
362 let raw = addr
363 .as_str()
364 .ok_or_else(|| TransportError::InvalidAddress(addr.to_string()))?;
365 let candidate = WebSocketConfig {
366 seed_urls: vec![raw.to_owned()],
367 ..self.config.clone()
368 };
369 candidate
370 .validate()
371 .map_err(TransportError::InvalidAddress)?;
372 if self.connection_state_sync(addr) == ConnectionState::Connected {
373 return Ok(());
374 }
375 {
376 let mut states = self
377 .runtime
378 .states
379 .lock()
380 .unwrap_or_else(|error| error.into_inner());
381 if matches!(states.get(addr), Some(ConnectionState::Connecting)) {
382 return Ok(());
383 }
384 states.insert(addr.clone(), ConnectionState::Connecting);
385 }
386 self.push_task(tokio::spawn(run_one_shot_dial(
387 self.runtime.clone(),
388 addr.clone(),
389 )));
390 Ok(())
391 }
392
393 pub fn connection_state_sync(&self, addr: &TransportAddr) -> ConnectionState {
394 if let Ok(pool) = self.runtime.pool.try_lock()
395 && pool.contains_key(addr)
396 {
397 return ConnectionState::Connected;
398 }
399 self.runtime
400 .states
401 .lock()
402 .unwrap_or_else(|error| error.into_inner())
403 .get(addr)
404 .cloned()
405 .unwrap_or(ConnectionState::None)
406 }
407
408 pub async fn close_connection_async(&self, addr: &TransportAddr) {
409 self.runtime.pool.lock().await.remove(addr);
410 self.runtime
411 .states
412 .lock()
413 .unwrap_or_else(|error| error.into_inner())
414 .remove(addr);
415 }
416
417 pub fn close_connection_detached(&self, addr: &TransportAddr) {
419 let runtime = self.runtime.clone();
420 let addr = addr.clone();
421 tokio::spawn(async move {
422 runtime.pool.lock().await.remove(&addr);
423 runtime
424 .states
425 .lock()
426 .unwrap_or_else(|error| error.into_inner())
427 .remove(&addr);
428 });
429 }
430
431 pub(crate) async fn restart_after_network_change(&mut self) -> Result<bool, TransportError> {
438 if !(self.state.is_operational() || self.state.can_start()) {
439 return Ok(false);
440 }
441 if self.state.is_operational() {
442 let mut pool = self.runtime.pool.lock().await;
444 self.runtime
445 .network_rebind_generation
446 .send_modify(|generation| *generation = generation.wrapping_add(1));
447 pool.clear();
448 self.runtime
449 .states
450 .lock()
451 .unwrap_or_else(|error| error.into_inner())
452 .clear();
453 drop(pool);
454 return Ok(true);
455 }
456 let previous_state = self.state;
457 let previous_local_addr = self.local_addr;
458 match self.start_async().await {
459 Ok(()) => Ok(true),
460 Err(error) => {
461 self.runtime.running.store(false, Ordering::Release);
462 self.local_addr = previous_local_addr;
463 self.state = previous_state;
464 Err(error)
465 }
466 }
467 }
468
469 pub(crate) async fn rollback_network_change_start(
470 &mut self,
471 previous_state: TransportState,
472 ) -> Result<(), TransportError> {
473 if self.state.is_operational() {
474 self.stop_async().await?;
475 }
476 self.state = previous_state;
477 Ok(())
478 }
479}
480
481impl Transport for WebSocketTransport {
482 fn transport_id(&self) -> TransportId {
483 self.transport_id
484 }
485
486 fn transport_type(&self) -> &TransportType {
487 &TransportType::WEBSOCKET
488 }
489
490 fn state(&self) -> TransportState {
491 self.state
492 }
493
494 fn mtu(&self) -> u16 {
495 self.config.mtu()
496 }
497
498 fn start(&mut self) -> Result<(), TransportError> {
499 Err(TransportError::NotSupported(
500 "use start_async() for WebSocket transport".into(),
501 ))
502 }
503
504 fn stop(&mut self) -> Result<(), TransportError> {
505 Err(TransportError::NotSupported(
506 "use stop_async() for WebSocket transport".into(),
507 ))
508 }
509
510 fn send(&self, _addr: &TransportAddr, _data: &[u8]) -> Result<(), TransportError> {
511 Err(TransportError::NotSupported(
512 "use send_async() for WebSocket transport".into(),
513 ))
514 }
515
516 fn discover(&self) -> Result<Vec<DiscoveredPeer>, TransportError> {
517 Ok(self
518 .runtime
519 .discoveries
520 .lock()
521 .unwrap_or_else(|error| error.into_inner())
522 .drain(..)
523 .collect())
524 }
525
526 fn auto_connect(&self) -> bool {
527 true
528 }
529
530 fn accept_connections(&self) -> bool {
531 self.config.accept_connections()
532 }
533}
534
535async fn run_accept_loop(runtime: Runtime, listener: TcpListener) {
536 while runtime.running.load(Ordering::Acquire) {
537 let Ok((stream, peer_addr)) = listener.accept().await else {
538 continue;
539 };
540 let Ok(inbound_permit) = runtime.inbound_slots.clone().try_acquire_owned() else {
541 runtime
542 .stats
543 .connections_rejected
544 .fetch_add(1, Ordering::Relaxed);
545 continue;
546 };
547 let Ok(total_permit) = runtime.total_slots.clone().try_acquire_owned() else {
548 runtime
549 .stats
550 .connections_rejected
551 .fetch_add(1, Ordering::Relaxed);
552 continue;
553 };
554 let accepted_runtime = runtime.clone();
555 let accepted_network_rebind_generation = *runtime.network_rebind_generation.borrow();
556 tokio::spawn(async move {
557 let _inbound_permit = inbound_permit;
558 let _total_permit = total_permit;
559 if let Err(error) = accept_connection(
560 accepted_runtime,
561 stream,
562 peer_addr,
563 accepted_network_rebind_generation,
564 )
565 .await
566 {
567 debug!(%peer_addr, %error, "WebSocket inbound connection ended");
568 }
569 });
570 }
571}
572
573#[allow(clippy::result_large_err)]
574async fn accept_connection(
575 runtime: Runtime,
576 stream: TcpStream,
577 peer_addr: SocketAddr,
578 network_rebind_generation: u64,
579) -> Result<(), TransportError> {
580 let path = runtime.config.path().to_owned();
581 let callback = move |request: &Request, response: Response| {
582 if request.uri().path() == path {
583 Ok(response)
584 } else {
585 let mut error = ErrorResponse::new(Some("not found".into()));
586 *error.status_mut() = tokio_tungstenite::tungstenite::http::StatusCode::NOT_FOUND;
587 Err(error)
588 }
589 };
590 let websocket =
591 accept_hdr_async_with_config(stream, callback, Some(runtime.websocket_config()))
592 .await
593 .map_err(|_| TransportError::ConnectionRefused)?;
594 let generation = runtime.next_generation();
595 let addr = TransportAddr::from_string(&format!("ws-peer://{peer_addr}/{generation}"));
596 run_connection(
597 runtime,
598 addr,
599 websocket,
600 generation,
601 Direction::Inbound,
602 false,
603 network_rebind_generation,
604 )
605 .await
606}
607
608async fn run_seed_dialer(runtime: Runtime, addr: TransportAddr) {
609 let mut delay_ms = runtime.config.reconnect_initial_ms();
610 let mut network_rebind_generation = runtime.network_rebind_generation.subscribe();
611 let mut observed_network_rebind = *network_rebind_generation.borrow_and_update();
612 while runtime.running.load(Ordering::Acquire) {
613 runtime.set_state(&addr, ConnectionState::Connecting);
614 runtime
615 .stats
616 .reconnect_attempts
617 .fetch_add(1, Ordering::Relaxed);
618 let result = tokio::select! {
619 result = dial_and_run(runtime.clone(), addr.clone()) => result,
620 changed = network_rebind_generation.changed() => {
621 if changed.is_err() {
622 break;
623 }
624 observed_network_rebind = *network_rebind_generation.borrow_and_update();
625 continue;
626 }
627 };
628 if !runtime.running.load(Ordering::Acquire) {
629 break;
630 }
631 match result {
632 Ok(()) => {
633 delay_ms = runtime.config.reconnect_initial_ms();
634 }
635 Err(error) => {
636 runtime.set_state(&addr, ConnectionState::Failed(error.to_string()));
637 debug!(remote_addr = %addr, %error, "WebSocket seed connection failed");
638 delay_ms = delay_ms
639 .saturating_mul(2)
640 .min(runtime.config.reconnect_max_ms());
641 }
642 }
643 let requested_network_rebind = *network_rebind_generation.borrow_and_update();
644 if requested_network_rebind != observed_network_rebind {
645 observed_network_rebind = requested_network_rebind;
646 continue;
648 }
649 tokio::select! {
650 _ = tokio::time::sleep(Duration::from_millis(delay_ms)) => {}
651 changed = network_rebind_generation.changed() => {
652 if changed.is_err() {
653 break;
654 }
655 observed_network_rebind = *network_rebind_generation.borrow_and_update();
656 }
657 }
658 }
659}
660
661async fn run_one_shot_dial(runtime: Runtime, addr: TransportAddr) {
662 let result = dial_and_run(runtime.clone(), addr.clone()).await;
663 if let Err(error) = result {
664 runtime.set_state(&addr, ConnectionState::Failed(error.to_string()));
665 }
666}
667
668async fn dial_and_run(runtime: Runtime, addr: TransportAddr) -> Result<(), TransportError> {
669 let network_rebind_generation = *runtime.network_rebind_generation.borrow();
670 let _slot = runtime
671 .total_slots
672 .clone()
673 .try_acquire_owned()
674 .map_err(|_| TransportError::ConnectionRefused)?;
675 let url = addr
676 .as_str()
677 .ok_or_else(|| TransportError::InvalidAddress(addr.to_string()))?;
678 let connect = connect_async_with_config(url, Some(runtime.websocket_config()), false);
679 let (websocket, _) = tokio::time::timeout(
680 Duration::from_millis(runtime.config.connect_timeout_ms()),
681 connect,
682 )
683 .await
684 .map_err(|_| TransportError::Timeout)?
685 .map_err(|error| TransportError::StartFailed(error.to_string()))?;
686 let generation = runtime.next_generation();
687 run_connection(
688 runtime,
689 addr,
690 websocket,
691 generation,
692 Direction::Outbound,
693 true,
694 network_rebind_generation,
695 )
696 .await
697}
698
699async fn run_connection<S>(
700 runtime: Runtime,
701 addr: TransportAddr,
702 websocket: WebSocketStream<S>,
703 generation: u64,
704 direction: Direction,
705 request_key_hint: bool,
706 expected_network_rebind_generation: u64,
707) -> Result<(), TransportError>
708where
709 S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
710{
711 let (mut sink, mut stream) = websocket.split();
712 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(runtime.config.max_send_queue());
713 {
714 let mut pool = runtime.pool.lock().await;
715 if *runtime.network_rebind_generation.borrow() != expected_network_rebind_generation {
717 return Err(TransportError::StartFailed(
718 "network changed during WebSocket handshake".into(),
719 ));
720 }
721 if pool.contains_key(&addr) {
722 return Err(TransportError::AlreadyStarted);
723 }
724 pool.insert(addr.clone(), Connection { generation, tx });
725 }
726 runtime.set_state(&addr, ConnectionState::Connected);
727 runtime
728 .stats
729 .connections_opened
730 .fetch_add(1, Ordering::Relaxed);
731
732 let mut pending_nonce = request_key_hint.then(|| rand::rng().random::<u64>());
733 if let Some(nonce) = pending_nonce {
734 sink.send(Message::Binary(
735 LocalKeyHint::Request { nonce }.encode().into(),
736 ))
737 .await
738 .map_err(|error| TransportError::SendFailed(error.to_string()))?;
739 }
740
741 let started = tokio::time::Instant::now();
742 let mut last_received = started;
743 let ping_secs = runtime.config.ping_interval_secs();
744 let idle_secs = runtime.config.idle_timeout_secs();
745 let mut ping = tokio::time::interval(if ping_secs == 0 {
746 Duration::from_secs(24 * 60 * 60)
747 } else {
748 Duration::from_secs(ping_secs)
749 });
750 ping.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
751 let mut check = tokio::time::interval(Duration::from_secs(1));
752 check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
753
754 enum Event {
755 Outbound(Option<Vec<u8>>),
756 Inbound(Option<Result<Message, tokio_tungstenite::tungstenite::Error>>),
757 Ping,
758 Check,
759 }
760
761 let outcome = loop {
762 let event = tokio::select! {
763 outbound = rx.recv() => Event::Outbound(outbound),
764 inbound = stream.next() => Event::Inbound(inbound),
765 _ = ping.tick(), if ping_secs > 0 => Event::Ping,
766 _ = check.tick() => Event::Check,
767 };
768 match event {
769 Event::Outbound(Some(data)) => {
770 let len = data.len();
771 if let Err(error) = sink.send(Message::Binary(data.into())).await {
772 break Err(TransportError::SendFailed(error.to_string()));
773 }
774 runtime.stats.frames_sent.fetch_add(1, Ordering::Relaxed);
775 runtime
776 .stats
777 .bytes_sent
778 .fetch_add(len as u64, Ordering::Relaxed);
779 }
780 Event::Outbound(None) => break Ok(()),
781 Event::Inbound(Some(Ok(Message::Binary(data)))) => {
782 last_received = tokio::time::Instant::now();
783 if let Some(hint) = LocalKeyHint::decode(&data) {
784 match hint {
785 LocalKeyHint::Request { nonce } => {
786 let reply = LocalKeyHint::Response {
787 nonce,
788 pubkey: runtime.local_pubkey,
789 };
790 if let Err(error) =
791 sink.send(Message::Binary(reply.encode().into())).await
792 {
793 break Err(TransportError::SendFailed(error.to_string()));
794 }
795 }
796 LocalKeyHint::Response { nonce, pubkey }
797 if pending_nonce == Some(nonce) =>
798 {
799 pending_nonce = None;
800 if pubkey != runtime.local_pubkey
801 && let Ok(pubkey) = XOnlyPublicKey::from_slice(&pubkey)
802 {
803 runtime
804 .discoveries
805 .lock()
806 .unwrap_or_else(|error| error.into_inner())
807 .push_back(DiscoveredPeer::with_hint(
808 runtime.transport_id,
809 addr.clone(),
810 pubkey,
811 ));
812 }
813 }
814 LocalKeyHint::Response { .. } => {}
815 }
816 continue;
817 }
818 if data.len() > runtime.config.max_frame_bytes()
819 || validate_websocket_record(&data).is_err()
820 {
821 runtime.stats.invalid_frames.fetch_add(1, Ordering::Relaxed);
822 break Err(TransportError::RecvFailed(
823 "invalid WebSocket FIPS physical record".into(),
824 ));
825 }
826 let len = data.len();
827 let packet = ReceivedPacket::with_timestamp(
828 runtime.transport_id,
829 addr.clone(),
830 PacketBuffer::new(data.to_vec()),
831 now_ms(),
832 );
833 if runtime.packet_tx.send(packet).is_err() {
834 break Err(TransportError::RecvFailed(
835 "node packet channel closed".into(),
836 ));
837 }
838 runtime
839 .stats
840 .frames_received
841 .fetch_add(1, Ordering::Relaxed);
842 runtime
843 .stats
844 .bytes_received
845 .fetch_add(len as u64, Ordering::Relaxed);
846 }
847 Event::Inbound(Some(Ok(Message::Ping(payload)))) => {
848 last_received = tokio::time::Instant::now();
849 if let Err(error) = sink.send(Message::Pong(payload)).await {
850 break Err(TransportError::SendFailed(error.to_string()));
851 }
852 }
853 Event::Inbound(Some(Ok(Message::Pong(_)))) => {
854 last_received = tokio::time::Instant::now();
855 }
856 Event::Inbound(Some(Ok(Message::Close(_))) | None) => break Ok(()),
857 Event::Inbound(Some(Ok(Message::Text(_) | Message::Frame(_)))) => {
858 runtime.stats.invalid_frames.fetch_add(1, Ordering::Relaxed);
859 break Err(TransportError::RecvFailed(
860 "WebSocket transport accepts binary messages only".into(),
861 ));
862 }
863 Event::Inbound(Some(Err(error))) => {
864 break Err(TransportError::RecvFailed(error.to_string()));
865 }
866 Event::Ping => {
867 if let Err(error) = sink.send(Message::Ping(Bytes::new())).await {
868 break Err(TransportError::SendFailed(error.to_string()));
869 }
870 }
871 Event::Check => {
872 if pending_nonce.is_some()
873 && started.elapsed()
874 >= Duration::from_millis(runtime.config.key_hint_timeout_ms())
875 {
876 break Err(TransportError::Timeout);
877 }
878 if idle_secs > 0 && last_received.elapsed() >= Duration::from_secs(idle_secs) {
879 break Err(TransportError::Timeout);
880 }
881 }
882 }
883 };
884
885 {
886 let mut pool = runtime.pool.lock().await;
887 if pool
888 .get(&addr)
889 .is_some_and(|connection| connection.generation == generation)
890 {
891 pool.remove(&addr);
892 }
893 }
894 runtime.clear_state_if(&addr, generation);
895 runtime
896 .stats
897 .connections_closed
898 .fetch_add(1, Ordering::Relaxed);
899 debug!(
900 transport_id = %runtime.transport_id,
901 remote_addr = %addr,
902 ?direction,
903 "WebSocket physical connection closed"
904 );
905 outcome
906}
907
908fn now_ms() -> u64 {
909 SystemTime::now()
910 .duration_since(UNIX_EPOCH)
911 .unwrap_or_default()
912 .as_millis() as u64
913}
914
915fn validate_websocket_record(data: &[u8]) -> Result<(), String> {
916 if validate_direct_fsp_transport_fragment(data) {
917 return Ok(());
918 }
919 validate_stream_record(data).map_err(|error| format!("invalid FIPS physical record: {error}"))
920}
921
922#[cfg(test)]
923mod tests;