Skip to main content

fips_core/transport/
websocket.rs

1//! Binary WebSocket physical transport.
2//!
3//! One binary WebSocket message carries one FIPS physical record. A tiny
4//! nonce-bound key-hint exchange precedes Noise IK when a client has only a
5//! seed URL; the hint is untrusted routing metadata and never bypasses FIPS
6//! identity authentication or ACLs.
7
8use 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    // Only explicit network changes advance this value; ordinary closes retain backoff.
110    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
150/// Generic WebSocket physical transport.
151pub 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            // The URL identifies the operator-configured physical dial even
226            // when simultaneous FIPS initiation makes this node the responder.
227            // Handshake role is not transport direction.
228            true
229        } else if !handshake_is_initiator {
230            // Accepting clients on a configured listener is an explicit
231            // operator choice; every promoted client has already completed
232            // the authenticated FIPS handshake.
233            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    /// Replace TCP-backed WebSocket streams after a confirmed network change.
418    ///
419    /// An established stream can remain locally "connected" long after its
420    /// source address or NAT mapping vanished. Dropping the connection senders
421    /// closes stale streams and wakes the existing seed dialers without
422    /// releasing and racing to reacquire a configured listener socket.
423    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            // Serialize the generation cutover with outbound pool insertion.
429            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            // The rebind itself closed this connection, so make one prompt replacement attempt.
633            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        // A handshake started on the previous underlay must not repopulate the cleared pool.
702        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;