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(crate) async fn restart_after_network_change(&mut self) -> Result<bool, TransportError> {
424 if !(self.state.is_operational() || self.state.can_start()) {
425 return Ok(false);
426 }
427 if self.state.is_operational() {
428 let mut pool = self.runtime.pool.lock().await;
430 self.runtime
431 .network_rebind_generation
432 .send_modify(|generation| *generation = generation.wrapping_add(1));
433 pool.clear();
434 self.runtime
435 .states
436 .lock()
437 .unwrap_or_else(|error| error.into_inner())
438 .clear();
439 drop(pool);
440 return Ok(true);
441 }
442 let previous_state = self.state;
443 let previous_local_addr = self.local_addr;
444 match self.start_async().await {
445 Ok(()) => Ok(true),
446 Err(error) => {
447 self.runtime.running.store(false, Ordering::Release);
448 self.local_addr = previous_local_addr;
449 self.state = previous_state;
450 Err(error)
451 }
452 }
453 }
454
455 pub(crate) async fn rollback_network_change_start(
456 &mut self,
457 previous_state: TransportState,
458 ) -> Result<(), TransportError> {
459 if self.state.is_operational() {
460 self.stop_async().await?;
461 }
462 self.state = previous_state;
463 Ok(())
464 }
465}
466
467impl Transport for WebSocketTransport {
468 fn transport_id(&self) -> TransportId {
469 self.transport_id
470 }
471
472 fn transport_type(&self) -> &TransportType {
473 &TransportType::WEBSOCKET
474 }
475
476 fn state(&self) -> TransportState {
477 self.state
478 }
479
480 fn mtu(&self) -> u16 {
481 self.config.mtu()
482 }
483
484 fn start(&mut self) -> Result<(), TransportError> {
485 Err(TransportError::NotSupported(
486 "use start_async() for WebSocket transport".into(),
487 ))
488 }
489
490 fn stop(&mut self) -> Result<(), TransportError> {
491 Err(TransportError::NotSupported(
492 "use stop_async() for WebSocket transport".into(),
493 ))
494 }
495
496 fn send(&self, _addr: &TransportAddr, _data: &[u8]) -> Result<(), TransportError> {
497 Err(TransportError::NotSupported(
498 "use send_async() for WebSocket transport".into(),
499 ))
500 }
501
502 fn discover(&self) -> Result<Vec<DiscoveredPeer>, TransportError> {
503 Ok(self
504 .runtime
505 .discoveries
506 .lock()
507 .unwrap_or_else(|error| error.into_inner())
508 .drain(..)
509 .collect())
510 }
511
512 fn auto_connect(&self) -> bool {
513 true
514 }
515
516 fn accept_connections(&self) -> bool {
517 self.config.accept_connections()
518 }
519}
520
521async fn run_accept_loop(runtime: Runtime, listener: TcpListener) {
522 while runtime.running.load(Ordering::Acquire) {
523 let Ok((stream, peer_addr)) = listener.accept().await else {
524 continue;
525 };
526 let Ok(inbound_permit) = runtime.inbound_slots.clone().try_acquire_owned() else {
527 runtime
528 .stats
529 .connections_rejected
530 .fetch_add(1, Ordering::Relaxed);
531 continue;
532 };
533 let Ok(total_permit) = runtime.total_slots.clone().try_acquire_owned() else {
534 runtime
535 .stats
536 .connections_rejected
537 .fetch_add(1, Ordering::Relaxed);
538 continue;
539 };
540 let accepted_runtime = runtime.clone();
541 let accepted_network_rebind_generation = *runtime.network_rebind_generation.borrow();
542 tokio::spawn(async move {
543 let _inbound_permit = inbound_permit;
544 let _total_permit = total_permit;
545 if let Err(error) = accept_connection(
546 accepted_runtime,
547 stream,
548 peer_addr,
549 accepted_network_rebind_generation,
550 )
551 .await
552 {
553 debug!(%peer_addr, %error, "WebSocket inbound connection ended");
554 }
555 });
556 }
557}
558
559#[allow(clippy::result_large_err)]
560async fn accept_connection(
561 runtime: Runtime,
562 stream: TcpStream,
563 peer_addr: SocketAddr,
564 network_rebind_generation: u64,
565) -> Result<(), TransportError> {
566 let path = runtime.config.path().to_owned();
567 let callback = move |request: &Request, response: Response| {
568 if request.uri().path() == path {
569 Ok(response)
570 } else {
571 let mut error = ErrorResponse::new(Some("not found".into()));
572 *error.status_mut() = tokio_tungstenite::tungstenite::http::StatusCode::NOT_FOUND;
573 Err(error)
574 }
575 };
576 let websocket =
577 accept_hdr_async_with_config(stream, callback, Some(runtime.websocket_config()))
578 .await
579 .map_err(|_| TransportError::ConnectionRefused)?;
580 let generation = runtime.next_generation();
581 let addr = TransportAddr::from_string(&format!("ws-peer://{peer_addr}/{generation}"));
582 run_connection(
583 runtime,
584 addr,
585 websocket,
586 generation,
587 Direction::Inbound,
588 false,
589 network_rebind_generation,
590 )
591 .await
592}
593
594async fn run_seed_dialer(runtime: Runtime, addr: TransportAddr) {
595 let mut delay_ms = runtime.config.reconnect_initial_ms();
596 let mut network_rebind_generation = runtime.network_rebind_generation.subscribe();
597 let mut observed_network_rebind = *network_rebind_generation.borrow_and_update();
598 while runtime.running.load(Ordering::Acquire) {
599 runtime.set_state(&addr, ConnectionState::Connecting);
600 runtime
601 .stats
602 .reconnect_attempts
603 .fetch_add(1, Ordering::Relaxed);
604 let result = tokio::select! {
605 result = dial_and_run(runtime.clone(), addr.clone()) => result,
606 changed = network_rebind_generation.changed() => {
607 if changed.is_err() {
608 break;
609 }
610 observed_network_rebind = *network_rebind_generation.borrow_and_update();
611 continue;
612 }
613 };
614 if !runtime.running.load(Ordering::Acquire) {
615 break;
616 }
617 match result {
618 Ok(()) => {
619 delay_ms = runtime.config.reconnect_initial_ms();
620 }
621 Err(error) => {
622 runtime.set_state(&addr, ConnectionState::Failed(error.to_string()));
623 debug!(remote_addr = %addr, %error, "WebSocket seed connection failed");
624 delay_ms = delay_ms
625 .saturating_mul(2)
626 .min(runtime.config.reconnect_max_ms());
627 }
628 }
629 let requested_network_rebind = *network_rebind_generation.borrow_and_update();
630 if requested_network_rebind != observed_network_rebind {
631 observed_network_rebind = requested_network_rebind;
632 continue;
634 }
635 tokio::select! {
636 _ = tokio::time::sleep(Duration::from_millis(delay_ms)) => {}
637 changed = network_rebind_generation.changed() => {
638 if changed.is_err() {
639 break;
640 }
641 observed_network_rebind = *network_rebind_generation.borrow_and_update();
642 }
643 }
644 }
645}
646
647async fn run_one_shot_dial(runtime: Runtime, addr: TransportAddr) {
648 let result = dial_and_run(runtime.clone(), addr.clone()).await;
649 if let Err(error) = result {
650 runtime.set_state(&addr, ConnectionState::Failed(error.to_string()));
651 }
652}
653
654async fn dial_and_run(runtime: Runtime, addr: TransportAddr) -> Result<(), TransportError> {
655 let network_rebind_generation = *runtime.network_rebind_generation.borrow();
656 let _slot = runtime
657 .total_slots
658 .clone()
659 .try_acquire_owned()
660 .map_err(|_| TransportError::ConnectionRefused)?;
661 let url = addr
662 .as_str()
663 .ok_or_else(|| TransportError::InvalidAddress(addr.to_string()))?;
664 let connect = connect_async_with_config(url, Some(runtime.websocket_config()), false);
665 let (websocket, _) = tokio::time::timeout(
666 Duration::from_millis(runtime.config.connect_timeout_ms()),
667 connect,
668 )
669 .await
670 .map_err(|_| TransportError::Timeout)?
671 .map_err(|error| TransportError::StartFailed(error.to_string()))?;
672 let generation = runtime.next_generation();
673 run_connection(
674 runtime,
675 addr,
676 websocket,
677 generation,
678 Direction::Outbound,
679 true,
680 network_rebind_generation,
681 )
682 .await
683}
684
685async fn run_connection<S>(
686 runtime: Runtime,
687 addr: TransportAddr,
688 websocket: WebSocketStream<S>,
689 generation: u64,
690 direction: Direction,
691 request_key_hint: bool,
692 expected_network_rebind_generation: u64,
693) -> Result<(), TransportError>
694where
695 S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
696{
697 let (mut sink, mut stream) = websocket.split();
698 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(runtime.config.max_send_queue());
699 {
700 let mut pool = runtime.pool.lock().await;
701 if *runtime.network_rebind_generation.borrow() != expected_network_rebind_generation {
703 return Err(TransportError::StartFailed(
704 "network changed during WebSocket handshake".into(),
705 ));
706 }
707 if pool.contains_key(&addr) {
708 return Err(TransportError::AlreadyStarted);
709 }
710 pool.insert(addr.clone(), Connection { generation, tx });
711 }
712 runtime.set_state(&addr, ConnectionState::Connected);
713 runtime
714 .stats
715 .connections_opened
716 .fetch_add(1, Ordering::Relaxed);
717
718 let mut pending_nonce = request_key_hint.then(|| rand::rng().random::<u64>());
719 if let Some(nonce) = pending_nonce {
720 sink.send(Message::Binary(
721 LocalKeyHint::Request { nonce }.encode().into(),
722 ))
723 .await
724 .map_err(|error| TransportError::SendFailed(error.to_string()))?;
725 }
726
727 let started = tokio::time::Instant::now();
728 let mut last_received = started;
729 let ping_secs = runtime.config.ping_interval_secs();
730 let idle_secs = runtime.config.idle_timeout_secs();
731 let mut ping = tokio::time::interval(if ping_secs == 0 {
732 Duration::from_secs(24 * 60 * 60)
733 } else {
734 Duration::from_secs(ping_secs)
735 });
736 ping.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
737 let mut check = tokio::time::interval(Duration::from_secs(1));
738 check.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
739
740 enum Event {
741 Outbound(Option<Vec<u8>>),
742 Inbound(Option<Result<Message, tokio_tungstenite::tungstenite::Error>>),
743 Ping,
744 Check,
745 }
746
747 let outcome = loop {
748 let event = tokio::select! {
749 outbound = rx.recv() => Event::Outbound(outbound),
750 inbound = stream.next() => Event::Inbound(inbound),
751 _ = ping.tick(), if ping_secs > 0 => Event::Ping,
752 _ = check.tick() => Event::Check,
753 };
754 match event {
755 Event::Outbound(Some(data)) => {
756 let len = data.len();
757 if let Err(error) = sink.send(Message::Binary(data.into())).await {
758 break Err(TransportError::SendFailed(error.to_string()));
759 }
760 runtime.stats.frames_sent.fetch_add(1, Ordering::Relaxed);
761 runtime
762 .stats
763 .bytes_sent
764 .fetch_add(len as u64, Ordering::Relaxed);
765 }
766 Event::Outbound(None) => break Ok(()),
767 Event::Inbound(Some(Ok(Message::Binary(data)))) => {
768 last_received = tokio::time::Instant::now();
769 if let Some(hint) = LocalKeyHint::decode(&data) {
770 match hint {
771 LocalKeyHint::Request { nonce } => {
772 let reply = LocalKeyHint::Response {
773 nonce,
774 pubkey: runtime.local_pubkey,
775 };
776 if let Err(error) =
777 sink.send(Message::Binary(reply.encode().into())).await
778 {
779 break Err(TransportError::SendFailed(error.to_string()));
780 }
781 }
782 LocalKeyHint::Response { nonce, pubkey }
783 if pending_nonce == Some(nonce) =>
784 {
785 pending_nonce = None;
786 if pubkey != runtime.local_pubkey
787 && let Ok(pubkey) = XOnlyPublicKey::from_slice(&pubkey)
788 {
789 runtime
790 .discoveries
791 .lock()
792 .unwrap_or_else(|error| error.into_inner())
793 .push_back(DiscoveredPeer::with_hint(
794 runtime.transport_id,
795 addr.clone(),
796 pubkey,
797 ));
798 }
799 }
800 LocalKeyHint::Response { .. } => {}
801 }
802 continue;
803 }
804 if data.len() > runtime.config.max_frame_bytes()
805 || validate_websocket_record(&data).is_err()
806 {
807 runtime.stats.invalid_frames.fetch_add(1, Ordering::Relaxed);
808 break Err(TransportError::RecvFailed(
809 "invalid WebSocket FIPS physical record".into(),
810 ));
811 }
812 let len = data.len();
813 let packet = ReceivedPacket::with_timestamp(
814 runtime.transport_id,
815 addr.clone(),
816 PacketBuffer::new(data.to_vec()),
817 now_ms(),
818 );
819 if runtime.packet_tx.send(packet).is_err() {
820 break Err(TransportError::RecvFailed(
821 "node packet channel closed".into(),
822 ));
823 }
824 runtime
825 .stats
826 .frames_received
827 .fetch_add(1, Ordering::Relaxed);
828 runtime
829 .stats
830 .bytes_received
831 .fetch_add(len as u64, Ordering::Relaxed);
832 }
833 Event::Inbound(Some(Ok(Message::Ping(payload)))) => {
834 last_received = tokio::time::Instant::now();
835 if let Err(error) = sink.send(Message::Pong(payload)).await {
836 break Err(TransportError::SendFailed(error.to_string()));
837 }
838 }
839 Event::Inbound(Some(Ok(Message::Pong(_)))) => {
840 last_received = tokio::time::Instant::now();
841 }
842 Event::Inbound(Some(Ok(Message::Close(_))) | None) => break Ok(()),
843 Event::Inbound(Some(Ok(Message::Text(_) | Message::Frame(_)))) => {
844 runtime.stats.invalid_frames.fetch_add(1, Ordering::Relaxed);
845 break Err(TransportError::RecvFailed(
846 "WebSocket transport accepts binary messages only".into(),
847 ));
848 }
849 Event::Inbound(Some(Err(error))) => {
850 break Err(TransportError::RecvFailed(error.to_string()));
851 }
852 Event::Ping => {
853 if let Err(error) = sink.send(Message::Ping(Bytes::new())).await {
854 break Err(TransportError::SendFailed(error.to_string()));
855 }
856 }
857 Event::Check => {
858 if pending_nonce.is_some()
859 && started.elapsed()
860 >= Duration::from_millis(runtime.config.key_hint_timeout_ms())
861 {
862 break Err(TransportError::Timeout);
863 }
864 if idle_secs > 0 && last_received.elapsed() >= Duration::from_secs(idle_secs) {
865 break Err(TransportError::Timeout);
866 }
867 }
868 }
869 };
870
871 {
872 let mut pool = runtime.pool.lock().await;
873 if pool
874 .get(&addr)
875 .is_some_and(|connection| connection.generation == generation)
876 {
877 pool.remove(&addr);
878 }
879 }
880 runtime.clear_state_if(&addr, generation);
881 runtime
882 .stats
883 .connections_closed
884 .fetch_add(1, Ordering::Relaxed);
885 debug!(
886 transport_id = %runtime.transport_id,
887 remote_addr = %addr,
888 ?direction,
889 "WebSocket physical connection closed"
890 );
891 outcome
892}
893
894fn now_ms() -> u64 {
895 SystemTime::now()
896 .duration_since(UNIX_EPOCH)
897 .unwrap_or_default()
898 .as_millis() as u64
899}
900
901fn validate_websocket_record(data: &[u8]) -> Result<(), String> {
902 if validate_direct_fsp_transport_fragment(data) {
903 return Ok(());
904 }
905 validate_stream_record(data).map_err(|error| format!("invalid FIPS physical record: {error}"))
906}
907
908#[cfg(test)]
909mod tests;