Skip to main content

nautilus_network/socket/
client.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Raw TCP client with optional TLS, suffix framing, heartbeats, and automatic reconnection.
17//!
18//! # State management
19//!
20//! The client tracks active, reconnecting, disconnecting, and closed states. State changes notify
21//! waiting tasks immediately instead of relying only on periodic checks.
22//!
23//! # Connection ownership
24//!
25//! A controller owns the connection lifecycle. A dedicated reader task passes complete messages to
26//! the configured callback, while a dedicated writer task serializes concurrent sends received over
27//! a channel.
28//!
29//! # Framing and heartbeats
30//!
31//! The configured suffix frames the byte stream in both directions. The writer appends it to sent
32//! messages and heartbeats, and the reader uses it to split incoming data into complete messages.
33//! Heartbeats are optional and run in a separate task.
34//!
35//! # Reconnection
36//!
37//! The writer buffers messages while reconnecting. A successful reconnect installs the replacement
38//! writer, drains that buffer, restarts the reader, and then invokes the configured
39//! post-reconnection callback.
40
41use std::{
42    collections::VecDeque,
43    fmt::Debug,
44    path::Path,
45    pin::pin,
46    sync::{
47        Arc,
48        atomic::{AtomicU8, Ordering},
49    },
50    time::Duration,
51};
52
53use bytes::Bytes;
54use nautilus_cryptography::providers::install_cryptographic_provider;
55use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
56use tokio_tungstenite::tungstenite::{Error, client::IntoClientRequest, stream::Mode};
57
58use super::{SocketConfig, TcpMessageHandler, TcpReader, TcpWriter, WriterCommand};
59use crate::{
60    SocketStateSink,
61    backoff::{
62        ExponentialBackoff, RECONNECT_STABILITY_THRESHOLD, ReconnectThrottle, wait_reconnect_delay,
63    },
64    dst,
65    error::{SendError, is_connection_drop_io_error},
66    logging::{log_task_aborted, log_task_started, log_task_stopped},
67    mode::{
68        ConnectionMode, ControllerLifecycle, ReadSessionFence, ReconnectOutcome,
69        ReconnectRequestOutcome,
70    },
71    net::TcpStream,
72    tls::{create_tls_config_from_certs_dir, tcp_tls},
73};
74
75// Connection timing constants
76const CONNECTION_STATE_CHECK_INTERVAL_MS: u64 = 10;
77const GRACEFUL_SHUTDOWN_DELAY_MS: u64 = 100;
78const GRACEFUL_SHUTDOWN_TIMEOUT_SECS: u64 = 5;
79const WRITE_TIMEOUT_SECS: u64 = 5;
80
81// Maximum buffer size for read operations (10 MB)
82const MAX_READ_BUFFER_BYTES: usize = 10 * 1024 * 1024;
83
84struct BufferedWrite {
85    data: Bytes,
86    replay_key: Option<u64>,
87}
88
89/// Produces protocol messages that must precede buffered application writes after reconnect.
90pub type SocketReconnectReplay = Arc<dyn Fn() -> Vec<Bytes> + Send + Sync>;
91
92struct SocketClientInner {
93    config: SocketConfig,
94    connector: Option<Arc<rustls::ClientConfig>>,
95    read_task: tokio::task::JoinHandle<()>,
96    read_fence: ReadSessionFence,
97    write_task: tokio::task::JoinHandle<()>,
98    writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
99    heartbeat_task: Option<tokio::task::JoinHandle<()>>,
100    connection_mode: Arc<AtomicU8>,
101    state_notify: Arc<tokio::sync::Notify>,
102    connect_timeout: Duration,
103    backoff: ExponentialBackoff,
104    reconnect_throttle: ReconnectThrottle,
105    reconnect_max_attempts: Option<u32>,
106    reconnect_attempt_count: u32,
107    state_sink: Option<SocketStateSink>,
108}
109
110impl SocketClientInner {
111    /// Connects to a URL with the specified configuration.
112    ///
113    /// # Errors
114    ///
115    /// Returns an error if connection fails or configuration is invalid.
116    async fn connect_url(
117        config: SocketConfig,
118        state_sink: Option<SocketStateSink>,
119    ) -> anyhow::Result<Self> {
120        install_cryptographic_provider();
121
122        // Validate suffix is non-empty to prevent panic in read loop (windows(0) panics)
123        if config.suffix.is_empty() {
124            anyhow::bail!("Socket suffix cannot be empty: suffix is required for message framing");
125        }
126
127        // Adapters build this config by struct literal, bypassing the builder, so this is the only
128        // place the field invariants are enforced for them.
129        config.validate()?;
130
131        let connect_timeout = Duration::from_millis(config.connect_timeout_ms.unwrap_or(10_000));
132        let reconnect_backoff = ExponentialBackoff::new(
133            Duration::from_millis(config.reconnect_delay_initial_ms.unwrap_or(2_000)),
134            Duration::from_millis(config.reconnect_delay_max_ms.unwrap_or(30_000)),
135            config.reconnect_backoff_factor.unwrap_or(1.5),
136            config.reconnect_jitter_ms.unwrap_or(100),
137            true, // immediate-first
138        )?;
139        let connector = if let Some(dir) = &config.certs_dir {
140            let config = create_tls_config_from_certs_dir(Path::new(dir), false)?;
141            Some(Arc::new(config))
142        } else {
143            None
144        };
145
146        // Retry initial connection with exponential backoff to handle transient DNS/network issues
147        let max_retries = config.connection_max_retries.unwrap_or(5);
148
149        let mut backoff = ExponentialBackoff::new(
150            Duration::from_millis(500),
151            Duration::from_secs(5),
152            2.0,
153            250,
154            false,
155        )?;
156
157        let mut attempt = 0;
158        let (reader, writer) = loop {
159            attempt += 1;
160
161            let last_error = match dst::time::timeout(
162                connect_timeout,
163                Self::tls_connect_with_server(&config.url, config.mode, connector.clone()),
164            )
165            .await
166            {
167                Ok(Ok(result)) => {
168                    if attempt > 1 {
169                        log::info!("Socket connection established after {attempt} attempts");
170                    }
171                    break result;
172                }
173                Ok(Err(e)) => {
174                    let error = e.to_string();
175                    log::warn!(
176                        "Socket connection attempt {attempt}/{max_retries} to {} failed: {error}",
177                        config.url,
178                    );
179                    error
180                }
181                Err(_) => {
182                    let error = format!(
183                        "Connection timeout after {:.1}s (possible DNS resolution failure)",
184                        connect_timeout.as_secs_f64()
185                    );
186                    log::warn!(
187                        "Socket connection attempt {attempt}/{max_retries} to {} timed out",
188                        config.url,
189                    );
190                    error
191                }
192            };
193
194            if attempt >= max_retries {
195                anyhow::bail!(
196                    "Failed to connect to {} after {} attempts: {}. \
197                    If this is a DNS error, check your network configuration and DNS settings.",
198                    config.url,
199                    max_retries,
200                    last_error,
201                );
202            }
203
204            let delay = backoff.next_duration();
205            log::debug!(
206                "Retrying in {delay:?} (attempt {}/{})",
207                attempt + 1,
208                max_retries
209            );
210            dst::time::sleep(delay).await;
211        };
212
213        log::debug!("Connected");
214
215        let connection_mode = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
216        let state_notify = Arc::new(tokio::sync::Notify::new());
217        let outcome =
218            ConnectionMode::complete_reconnect_with_sink(&connection_mode, state_sink.as_ref());
219        debug_assert_eq!(outcome, ReconnectOutcome::Reconnected);
220        let read_fence = ReadSessionFence::new();
221
222        let read_task = Self::spawn_read_task(
223            connection_mode.clone(),
224            read_fence.clone(),
225            reader,
226            config.message_handler.clone(),
227            config.suffix.clone(),
228            config.resolved_heartbeat_timeout(),
229        );
230
231        let (writer_tx, writer_rx) = tokio::sync::mpsc::unbounded_channel::<WriterCommand>();
232
233        let write_task = Self::spawn_write_task(
234            connection_mode.clone(),
235            state_notify.clone(),
236            writer,
237            writer_rx,
238            config.suffix.clone(),
239            state_sink.clone(),
240        );
241
242        let heartbeat_task = config.heartbeat.as_ref().map(|heartbeat| {
243            Self::spawn_heartbeat_task(
244                connection_mode.clone(),
245                heartbeat.interval_secs,
246                heartbeat.payload.clone(),
247                writer_tx.clone(),
248            )
249        });
250        let reconnect_max_attempts = config.reconnect_max_attempts;
251
252        Ok(Self {
253            config,
254            connector,
255            read_task,
256            read_fence,
257            write_task,
258            writer_tx,
259            heartbeat_task,
260            connection_mode,
261            state_notify,
262            connect_timeout,
263            backoff: reconnect_backoff,
264            reconnect_throttle: ReconnectThrottle::default(),
265            reconnect_max_attempts,
266            reconnect_attempt_count: 0,
267            state_sink,
268        })
269    }
270
271    /// Parses a URL into its socket address and request URL.
272    ///
273    /// Accepts either:
274    /// - Raw socket address: "host:port" → returns ("host:port", "scheme://host:port")
275    /// - Full URL: "scheme://host:port/path" → returns ("host:port", original URL)
276    ///
277    /// # Errors
278    ///
279    /// Returns an error if the URL is invalid or missing required components.
280    fn parse_socket_url(url: &str, mode: Mode) -> Result<(String, String), Error> {
281        if url.contains("://") {
282            // URL with scheme (e.g., "wss://host:port/path")
283            let parsed = url.parse::<http::Uri>().map_err(|e| {
284                Error::Io(std::io::Error::new(
285                    std::io::ErrorKind::InvalidInput,
286                    format!("Invalid URL: {e}"),
287                ))
288            })?;
289
290            let host = parsed.host().ok_or_else(|| {
291                Error::Io(std::io::Error::new(
292                    std::io::ErrorKind::InvalidInput,
293                    "URL missing host",
294                ))
295            })?;
296
297            let port = parsed
298                .port_u16()
299                .unwrap_or_else(|| match parsed.scheme_str() {
300                    Some("wss" | "https") => 443,
301                    Some("ws" | "http") => 80,
302                    _ => match mode {
303                        Mode::Tls => 443,
304                        Mode::Plain => 80,
305                    },
306                });
307
308            Ok((format!("{host}:{port}"), url.to_string()))
309        } else {
310            // Raw socket address (e.g., "host:port")
311            // Construct a proper URL for the request based on mode
312            let scheme = match mode {
313                Mode::Tls => "wss",
314                Mode::Plain => "ws",
315            };
316            Ok((url.to_string(), format!("{scheme}://{url}")))
317        }
318    }
319
320    /// Establish a TLS or plain TCP connection with the server.
321    ///
322    /// Accepts either a raw socket address (e.g., "host:port") or a full URL with scheme
323    /// (e.g., "wss://host:port"). For FIX/raw socket connections, use the host:port format.
324    /// For WebSocket-style connections, include the scheme.
325    ///
326    /// # Errors
327    ///
328    /// Returns an error if the connection cannot be established.
329    pub(crate) async fn tls_connect_with_server(
330        url: &str,
331        mode: Mode,
332        connector: Option<Arc<rustls::ClientConfig>>,
333    ) -> Result<(TcpReader, TcpWriter), Error> {
334        log::debug!("Connecting to {url}");
335
336        let (socket_addr, request_url) = Self::parse_socket_url(url, mode)?;
337        let tcp_result = TcpStream::connect(&socket_addr).await;
338
339        match tcp_result {
340            Ok(stream) => {
341                log::debug!("TCP connection established to {socket_addr}, proceeding with TLS");
342
343                crate::net::apply_socket_options(&stream);
344
345                let request = request_url.into_client_request()?;
346                tcp_tls(&request, mode, stream, connector)
347                    .await
348                    .map(tokio::io::split)
349            }
350            Err(e) => {
351                log::warn!("TCP connection failed to {socket_addr}: {e:?}");
352                Err(Error::Io(e))
353            }
354        }
355    }
356
357    /// Reconnects to the server.
358    ///
359    /// Makes a new connection with server, uses the new read and write halves
360    /// to update the reader and writer.
361    ///
362    /// The reconnect timeout bounds only connection establishment. Once the
363    /// new writer is handed to the writer task the swap runs to completion,
364    /// so buffered messages can never drain into a connection that lost its
365    /// reader to a timeout; the writer task bounds both the old-writer
366    /// shutdown and the buffer drain with its graceful-shutdown timeout.
367    async fn reconnect(
368        &mut self,
369        reconnect_replay: Option<&SocketReconnectReplay>,
370    ) -> Result<ReconnectOutcome, Error> {
371        log::info!("Reconnecting");
372
373        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
374            log::debug!("Reconnect aborted due to disconnect state");
375            return Ok(ReconnectOutcome::Aborted);
376        }
377
378        // Bound only connection establishment; the swap below must run to completion
379        let (reader, new_writer) = dst::time::timeout(
380            self.connect_timeout,
381            Self::tls_connect_with_server(
382                &self.config.url,
383                self.config.mode,
384                self.connector.clone(),
385            ),
386        )
387        .await
388        .map_err(|_| {
389            Error::Io(std::io::Error::new(
390                std::io::ErrorKind::TimedOut,
391                format!(
392                    "reconnection timed out after {}s",
393                    self.connect_timeout.as_secs_f64()
394                ),
395            ))
396        })??;
397
398        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
399            log::debug!("Reconnect aborted mid-flight (after connect)");
400            return Ok(ReconnectOutcome::Aborted);
401        }
402        log::debug!("Connected");
403
404        // Use a oneshot channel to synchronize with the writer task.
405        // We must verify that the buffer was successfully drained before transitioning to ACTIVE
406        // to prevent silent message loss if the new connection drops immediately.
407        let (tx, rx) = tokio::sync::oneshot::channel();
408        let command = if let Some(reconnect_replay) = reconnect_replay {
409            WriterCommand::UpdateWithReplay(new_writer, reconnect_replay(), tx)
410        } else {
411            WriterCommand::Update(new_writer, tx)
412        };
413
414        if let Err(e) = self.writer_tx.send(command) {
415            log::error!("{e}");
416            return Err(Error::Io(std::io::Error::new(
417                std::io::ErrorKind::BrokenPipe,
418                format!("Failed to send update command: {e}"),
419            )));
420        }
421
422        // Wait for writer to confirm it has drained the buffer
423        match rx.await {
424            Ok(true) => log::debug!("Writer confirmed buffer drain success"),
425            Ok(false) => {
426                log::warn!("Writer failed to drain buffer, aborting reconnect");
427                // Return error to trigger retry logic in controller
428                return Err(Error::Io(std::io::Error::other(
429                    "Failed to drain reconnection buffer",
430                )));
431            }
432            Err(e) => {
433                log::error!("Writer dropped update channel: {e}");
434                return Err(Error::Io(std::io::Error::new(
435                    std::io::ErrorKind::BrokenPipe,
436                    "Writer task dropped response channel",
437                )));
438            }
439        }
440
441        // Delay before closing connection
442        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
443
444        if ConnectionMode::from_atomic(&self.connection_mode).is_disconnect() {
445            log::debug!("Reconnect aborted mid-flight (after delay)");
446            return Ok(ReconnectOutcome::Aborted);
447        }
448
449        self.read_fence.invalidate();
450
451        if !self.read_task.is_finished() {
452            self.read_task.abort();
453            log_task_aborted("read");
454        }
455
456        // Atomically transition from Reconnect to Active
457        // This prevents race condition where disconnect could be requested between check and store
458        if ConnectionMode::complete_reconnect_with_sink(
459            &self.connection_mode,
460            self.state_sink.as_ref(),
461        ) == ReconnectOutcome::Aborted
462        {
463            log::debug!("Reconnect aborted (state changed during reconnect)");
464            return Ok(ReconnectOutcome::Aborted);
465        }
466
467        // Spawn new read task
468        self.read_fence = ReadSessionFence::new();
469        self.read_task = Self::spawn_read_task(
470            self.connection_mode.clone(),
471            self.read_fence.clone(),
472            reader,
473            self.config.message_handler.clone(),
474            self.config.suffix.clone(),
475            self.config.resolved_heartbeat_timeout(),
476        );
477
478        log::info!("Reconnect succeeded");
479        Ok(ReconnectOutcome::Reconnected)
480    }
481
482    /// Returns whether the read and write tasks are still running.
483    ///
484    /// Returns `true` if both the read and write tasks are still running.
485    /// There may be some delay between the connection closing and the
486    /// client detecting it.
487    #[inline]
488    #[must_use]
489    pub(crate) fn is_alive(&self) -> bool {
490        !self.read_task.is_finished() && !self.write_task.is_finished()
491    }
492
493    #[must_use]
494    fn spawn_read_task<R>(
495        connection_state: Arc<AtomicU8>,
496        read_fence: ReadSessionFence,
497        reader: R,
498        handler: Option<TcpMessageHandler>,
499        suffix: Vec<u8>,
500        heartbeat_timeout_secs: Option<u64>,
501    ) -> tokio::task::JoinHandle<()>
502    where
503        R: AsyncRead + Unpin + Send + 'static,
504    {
505        log_task_started("read");
506
507        // Interval between checking the connection mode
508        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
509        let heartbeat_timeout = heartbeat_timeout_secs.map(Duration::from_secs);
510
511        tokio::task::spawn(Self::run_read_loop(
512            connection_state,
513            read_fence,
514            reader,
515            handler,
516            suffix,
517            heartbeat_timeout,
518            check_interval,
519        ))
520    }
521
522    async fn run_read_loop<R>(
523        connection_state: Arc<AtomicU8>,
524        read_fence: ReadSessionFence,
525        mut reader: R,
526        handler: Option<TcpMessageHandler>,
527        suffix: Vec<u8>,
528        heartbeat_timeout: Option<Duration>,
529        check_interval: Duration,
530    ) where
531        R: AsyncRead + Unpin,
532    {
533        let mut buf = Vec::new();
534        let mut last_data_time = dst::time::Instant::now();
535
536        'read: loop {
537            if !ConnectionMode::from_atomic(&connection_state).is_active() || !read_fence.is_valid()
538            {
539                if !buf.is_empty() {
540                    log::debug!(
541                        "Dropping {} buffered bytes after socket session ended",
542                        buf.len()
543                    );
544                    buf.clear();
545                }
546                break;
547            }
548
549            match dst::time::timeout(check_interval, reader.read_buf(&mut buf)).await {
550                // Connection has been terminated or vector buffer is complete
551                Ok(Ok(0)) => {
552                    log::debug!("Connection closed by server");
553                    break;
554                }
555                Ok(Err(e)) => {
556                    log::debug!("Connection ended: {e}");
557                    break;
558                }
559                // Received bytes of data
560                Ok(Ok(bytes)) => {
561                    log::trace!("Received <binary> {bytes} bytes");
562                    last_data_time = dst::time::Instant::now();
563
564                    if !ConnectionMode::from_atomic(&connection_state).is_active()
565                        || !read_fence.is_valid()
566                    {
567                        log::debug!(
568                            "Dropping {} buffered bytes after socket session ended",
569                            buf.len()
570                        );
571                        buf.clear();
572                        break;
573                    }
574
575                    while let Some((i, _)) = &buf
576                        .windows(suffix.len())
577                        .enumerate()
578                        .find(|(_, pair)| pair.eq(&suffix))
579                    {
580                        let mut data: Vec<u8> = buf.drain(0..i + suffix.len()).collect();
581                        data.truncate(data.len() - suffix.len());
582
583                        if let Some(ref handler) = handler {
584                            if !ConnectionMode::from_atomic(&connection_state).is_active()
585                                || !read_fence.is_valid()
586                            {
587                                log::debug!(
588                                    "Dropping {} buffered bytes after socket session ended",
589                                    data.len() + buf.len()
590                                );
591                                buf.clear();
592                                break 'read;
593                            }
594                            handler(&data);
595                        }
596                    }
597
598                    if buf.len() > MAX_READ_BUFFER_BYTES {
599                        log::error!(
600                            "Read buffer exceeded maximum size ({MAX_READ_BUFFER_BYTES} bytes), closing connection"
601                        );
602                        break;
603                    }
604                }
605                Err(_) => {
606                    if let Some(timeout) = heartbeat_timeout {
607                        let silent_for = last_data_time.elapsed();
608                        if silent_for >= timeout {
609                            log::warn!(
610                                "Heartbeat timeout: no bytes received for {:.1}s",
611                                silent_for.as_secs_f64()
612                            );
613                            break;
614                        }
615                    }
616                }
617            }
618        }
619
620        log_task_stopped("read");
621    }
622
623    /// Drains buffered messages after reconnection completes.
624    ///
625    /// Attempts to send all buffered messages that were queued during reconnection.
626    /// Uses a peek-and-pop pattern to preserve messages if sending fails midway through the buffer.
627    ///
628    /// # Returns
629    ///
630    /// Returns `true` if a send error occurred (buffer may still contain unsent messages),
631    /// `false` if all messages were sent successfully (buffer is empty).
632    async fn drain_reconnect_buffer<W>(
633        buffer: &mut VecDeque<BufferedWrite>,
634        writer: &mut W,
635        suffix: &[u8],
636        replay: &[Bytes],
637    ) -> bool
638    where
639        W: AsyncWrite + Unpin,
640    {
641        if buffer.is_empty() {
642            return false;
643        }
644
645        let initial_buffer_len = buffer.len();
646        log::info!("Sending {initial_buffer_len} buffered messages after reconnection");
647
648        while let Some(buffered) = buffer.front() {
649            if buffered.replay_key.is_some()
650                && replay.iter().any(|replayed| replayed == &buffered.data)
651            {
652                buffer.pop_front();
653                continue;
654            }
655
656            let mut combined_msg = Vec::with_capacity(buffered.data.len() + suffix.len());
657            combined_msg.extend_from_slice(&buffered.data);
658            combined_msg.extend_from_slice(suffix);
659
660            if let Err(e) = writer.write_all(&combined_msg).await {
661                if is_connection_drop_io_error(&e) {
662                    log::warn!(
663                        "Failed to send buffered message with suffix after reconnection: {e}, {} messages remain in buffer",
664                        buffer.len()
665                    );
666                } else {
667                    log::error!(
668                        "Failed to send buffered message with suffix after reconnection: {e}, {} messages remain in buffer",
669                        buffer.len()
670                    );
671                }
672                return true;
673            }
674
675            buffer.pop_front();
676        }
677
678        log::info!("Successfully sent all {initial_buffer_len} buffered messages");
679
680        false
681    }
682
683    async fn replace_writer<W>(
684        active_writer: &mut W,
685        new_writer: W,
686        replay: Vec<Bytes>,
687        reconnect_buffer: &mut VecDeque<BufferedWrite>,
688        suffix: &[u8],
689    ) -> bool
690    where
691        W: AsyncWrite + Unpin,
692    {
693        log::debug!("Received new writer");
694        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
695
696        _ = dst::time::timeout(
697            Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
698            active_writer.shutdown(),
699        )
700        .await;
701
702        *active_writer = new_writer;
703        log::debug!("Updated writer");
704
705        let drain_result =
706            dst::time::timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS), async {
707                for replay_msg in &replay {
708                    let mut framed = Vec::with_capacity(replay_msg.len() + suffix.len());
709                    framed.extend_from_slice(replay_msg);
710                    framed.extend_from_slice(suffix);
711                    if let Err(e) = active_writer.write_all(&framed).await {
712                        log::warn!("Failed to send reconnect replay: {e}");
713                        return true;
714                    }
715                }
716
717                Self::drain_reconnect_buffer(reconnect_buffer, active_writer, suffix, &replay).await
718            })
719            .await;
720
721        let send_error = drain_result.unwrap_or_else(|_| {
722            log::warn!(
723                "Timed out sending reconnect replay and buffered messages, {} buffered messages remain",
724                reconnect_buffer.len()
725            );
726            true
727        });
728        !send_error
729    }
730
731    fn spawn_write_task<W>(
732        connection_state: Arc<AtomicU8>,
733        state_notify: Arc<tokio::sync::Notify>,
734        writer: W,
735        mut writer_rx: tokio::sync::mpsc::UnboundedReceiver<WriterCommand<W>>,
736        suffix: Vec<u8>,
737        state_sink: Option<SocketStateSink>,
738    ) -> tokio::task::JoinHandle<()>
739    where
740        W: AsyncWrite + Unpin + Send + 'static,
741    {
742        log_task_started("write");
743
744        // Interval between checking the connection mode
745        let check_interval = Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS);
746
747        tokio::task::spawn(async move {
748            let mut active_writer = writer;
749            let mut reconnect_buffer: VecDeque<BufferedWrite> = VecDeque::new();
750            let mut write_buf: Vec<u8> = Vec::new();
751
752            loop {
753                let mode = ConnectionMode::from_atomic(&connection_state);
754                if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
755                    break;
756                }
757
758                if mode.is_active() && !reconnect_buffer.is_empty() {
759                    let drain_result = dst::time::timeout(
760                        Duration::from_secs(WRITE_TIMEOUT_SECS),
761                        Self::drain_reconnect_buffer(
762                            &mut reconnect_buffer,
763                            &mut active_writer,
764                            &suffix,
765                            &[],
766                        ),
767                    )
768                    .await;
769                    let send_error = drain_result.unwrap_or_else(|_| {
770                        log::warn!(
771                            "Timed out draining reconnect buffer after {WRITE_TIMEOUT_SECS}s, {} messages remain",
772                            reconnect_buffer.len()
773                        );
774                        true
775                    });
776
777                    if send_error
778                        && ConnectionMode::request_reconnect_with_sink(
779                            &connection_state,
780                            state_sink.as_ref(),
781                        )
782                    {
783                        log::warn!("Writer triggering reconnect");
784                        state_notify.notify_one();
785                    }
786                    continue;
787                }
788
789                match dst::time::timeout(check_interval, writer_rx.recv()).await {
790                    Ok(Some(msg)) => {
791                        // Re-check connection mode after receiving a message
792                        let mode = ConnectionMode::from_atomic(&connection_state);
793                        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
794                            break;
795                        }
796
797                        match msg {
798                            WriterCommand::Update(new_writer, tx) => {
799                                let sent = Self::replace_writer(
800                                    &mut active_writer,
801                                    new_writer,
802                                    Vec::new(),
803                                    &mut reconnect_buffer,
804                                    &suffix,
805                                )
806                                .await;
807
808                                if let Err(e) = tx.send(sent) {
809                                    log::error!(
810                                        "Failed to report drain status to controller: {e:?}"
811                                    );
812                                }
813                            }
814                            WriterCommand::UpdateWithReplay(new_writer, replay, tx) => {
815                                let sent = Self::replace_writer(
816                                    &mut active_writer,
817                                    new_writer,
818                                    replay,
819                                    &mut reconnect_buffer,
820                                    &suffix,
821                                )
822                                .await;
823
824                                if let Err(e) = tx.send(sent) {
825                                    log::error!(
826                                        "Failed to report drain status to controller: {e:?}"
827                                    );
828                                }
829                            }
830                            WriterCommand::Send(data)
831                                if mode.is_reconnect() || !reconnect_buffer.is_empty() =>
832                            {
833                                log::debug!(
834                                    "Buffering message until reconnect drain completes ({} bytes)",
835                                    data.len()
836                                );
837                                Self::buffer_reconnect_write(&mut reconnect_buffer, data, None);
838                            }
839                            WriterCommand::SendOrReplay { key, data } if mode.is_reconnect() => {
840                                log::debug!(
841                                    "Buffering replayable message while reconnecting ({} bytes)",
842                                    data.len()
843                                );
844                                Self::buffer_reconnect_write(
845                                    &mut reconnect_buffer,
846                                    data,
847                                    Some(key),
848                                );
849                            }
850                            command @ (WriterCommand::Send(_)
851                            | WriterCommand::SendOrReplay { .. }) => {
852                                let (msg, replay_key) = match command {
853                                    WriterCommand::Send(data) => (data, None),
854                                    WriterCommand::SendOrReplay { key, data } => (data, Some(key)),
855                                    _ => unreachable!(),
856                                };
857                                write_buf.clear();
858                                write_buf.extend_from_slice(&msg);
859                                write_buf.extend_from_slice(&suffix);
860
861                                let write_result = dst::time::timeout(
862                                    Duration::from_secs(WRITE_TIMEOUT_SECS),
863                                    active_writer.write_all(&write_buf),
864                                )
865                                .await;
866                                let write_failed = match write_result {
867                                    Ok(Ok(())) => false,
868                                    Ok(Err(e)) => {
869                                        if is_connection_drop_io_error(&e) {
870                                            log::warn!("Failed to send message: {e}");
871                                        } else {
872                                            log::error!("Failed to send message: {e}");
873                                        }
874                                        true
875                                    }
876                                    Err(_) => {
877                                        log::warn!(
878                                            "Timed out sending message after {WRITE_TIMEOUT_SECS}s"
879                                        );
880                                        true
881                                    }
882                                };
883
884                                if write_failed {
885                                    Self::buffer_reconnect_write(
886                                        &mut reconnect_buffer,
887                                        msg,
888                                        replay_key,
889                                    );
890
891                                    // CAS: a disconnect landing mid-write must not be overwritten
892                                    if ConnectionMode::request_reconnect_with_sink(
893                                        &connection_state,
894                                        state_sink.as_ref(),
895                                    ) {
896                                        log::warn!("Writer triggering reconnect");
897                                        state_notify.notify_one();
898                                    }
899                                }
900                            }
901                        }
902                    }
903                    Ok(None) => {
904                        // Channel closed - writer task should terminate
905                        log::debug!("Writer channel closed, terminating writer task");
906                        break;
907                    }
908                    Err(_) => {
909                        // Timeout - just continue the loop
910                    }
911                }
912            }
913
914            // Attempt to shutdown the writer gracefully before exiting,
915            // we ignore any error as the writer may already be closed.
916            _ = dst::time::timeout(
917                Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS),
918                active_writer.shutdown(),
919            )
920            .await;
921
922            log_task_stopped("write");
923        })
924    }
925
926    fn buffer_reconnect_write(
927        buffer: &mut VecDeque<BufferedWrite>,
928        data: Bytes,
929        replay_key: Option<u64>,
930    ) {
931        if let Some(key) = replay_key {
932            buffer.retain(|buffered| buffered.replay_key != Some(key));
933        }
934        buffer.push_back(BufferedWrite { data, replay_key });
935    }
936
937    fn spawn_heartbeat_task(
938        connection_state: Arc<AtomicU8>,
939        interval_secs: u64,
940        message: Vec<u8>,
941        writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
942    ) -> tokio::task::JoinHandle<()> {
943        log_task_started("heartbeat");
944
945        tokio::task::spawn(async move {
946            let interval = Duration::from_secs(interval_secs);
947
948            loop {
949                dst::time::sleep(interval).await;
950
951                match ConnectionMode::from_u8(connection_state.load(Ordering::SeqCst)) {
952                    ConnectionMode::Active => {
953                        let msg = WriterCommand::Send(message.clone().into());
954
955                        match writer_tx.send(msg) {
956                            Ok(()) => log::trace!("Sent heartbeat to writer task"),
957                            Err(e) => {
958                                log::error!("Failed to send heartbeat to writer task: {e}");
959                            }
960                        }
961                    }
962                    ConnectionMode::Reconnect => {}
963                    ConnectionMode::Disconnect | ConnectionMode::Closed => break,
964                }
965            }
966
967            log_task_stopped("heartbeat");
968        })
969    }
970}
971
972impl Drop for SocketClientInner {
973    fn drop(&mut self) {
974        self.read_fence.invalidate();
975
976        if !self.read_task.is_finished() {
977            self.read_task.abort();
978            log_task_aborted("read");
979        }
980
981        if !self.write_task.is_finished() {
982            self.write_task.abort();
983            log_task_aborted("write");
984        }
985
986        if let Some(ref handle) = self.heartbeat_task.take()
987            && !handle.is_finished()
988        {
989            handle.abort();
990            log_task_aborted("heartbeat");
991        }
992    }
993}
994
995/// A suffix-framed TCP client with optional TLS and automatic reconnection.
996///
997/// The internal writer task serializes concurrent calls to [`Self::send_bytes`]. The configured
998/// suffix frames all sent and received messages, and an optional heartbeat task sends its payload
999/// at the configured interval. See [`SocketConfig`] for framing and reconnect policy.
1000pub struct SocketClient {
1001    pub(crate) controller_task: tokio::task::JoinHandle<()>,
1002    pub(crate) connection_mode: Arc<AtomicU8>,
1003    pub(crate) state_notify: Arc<tokio::sync::Notify>,
1004    pub(crate) connect_timeout: Duration,
1005    pub writer_tx: tokio::sync::mpsc::UnboundedSender<WriterCommand>,
1006    state_sink: Option<SocketStateSink>,
1007    controller_lifecycle: Arc<ControllerLifecycle>,
1008    controller_notify: Arc<tokio::sync::Notify>,
1009}
1010
1011/// Cloneable controller handle for requesting one raw socket reconnect.
1012#[derive(Clone)]
1013pub struct SocketReconnectHandle {
1014    connection_mode: Arc<AtomicU8>,
1015    state_sink: Option<SocketStateSink>,
1016    controller_lifecycle: Arc<ControllerLifecycle>,
1017    controller_notify: Arc<tokio::sync::Notify>,
1018}
1019
1020impl Debug for SocketReconnectHandle {
1021    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1022        f.debug_struct(stringify!(SocketReconnectHandle))
1023            .field(
1024                "connection_mode",
1025                &ConnectionMode::from_atomic(&self.connection_mode),
1026            )
1027            .finish_non_exhaustive()
1028    }
1029}
1030
1031impl SocketReconnectHandle {
1032    /// Requests that the controller replace the active transport.
1033    ///
1034    /// An accepted request reports the transport unavailable and wakes the controller. Rejected
1035    /// requests leave the transport state unchanged.
1036    #[must_use]
1037    pub fn request_reconnect(&self) -> ReconnectRequestOutcome {
1038        let Some(_request) = self.controller_lifecycle.enter_request() else {
1039            return ReconnectRequestOutcome::Closed;
1040        };
1041
1042        let outcome = ConnectionMode::request_reconnect_outcome_with_sink(
1043            &self.connection_mode,
1044            self.state_sink.as_ref(),
1045        );
1046
1047        if outcome == ReconnectRequestOutcome::Accepted {
1048            self.controller_notify.notify_one();
1049        }
1050        outcome
1051    }
1052}
1053
1054impl Debug for SocketClient {
1055    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1056        f.debug_struct(stringify!(SocketClient)).finish()
1057    }
1058}
1059
1060#[bon::bon]
1061impl SocketClient {
1062    /// Returns a builder for a new [`SocketClient`] connection.
1063    ///
1064    /// After a successful reconnect, `post_reconnection` runs after the replacement writer is
1065    /// installed, buffered messages are drained, and the replacement reader is started. The
1066    /// callback does not run after the initial connection.
1067    ///
1068    /// `state_sink` receives connection state changes. `reconnect_replay` produces protocol setup
1069    /// messages that the client sends before buffered application messages after a reconnect.
1070    ///
1071    /// # Errors
1072    ///
1073    /// Returns any error connecting to the server.
1074    #[builder(finish_fn = connect)]
1075    pub async fn builder(
1076        config: SocketConfig,
1077        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1078        state_sink: Option<SocketStateSink>,
1079        reconnect_replay: Option<SocketReconnectReplay>,
1080    ) -> anyhow::Result<Self> {
1081        let inner = SocketClientInner::connect_url(config, state_sink).await?;
1082        let writer_tx = inner.writer_tx.clone();
1083        let connection_mode = inner.connection_mode.clone();
1084        let state_notify = inner.state_notify.clone();
1085        let connect_timeout = inner.connect_timeout;
1086        let state_sink = inner.state_sink.clone();
1087        let controller_lifecycle = Arc::new(ControllerLifecycle::new());
1088        let controller_notify = Arc::new(tokio::sync::Notify::new());
1089        let controller_task = Self::spawn_controller_task(
1090            inner,
1091            connection_mode.clone(),
1092            state_notify.clone(),
1093            Arc::clone(&controller_lifecycle),
1094            Arc::clone(&controller_notify),
1095            post_reconnection,
1096            reconnect_replay,
1097        );
1098        controller_lifecycle.set_abort_handle(controller_task.abort_handle());
1099
1100        Ok(Self {
1101            controller_task,
1102            connection_mode,
1103            state_notify,
1104            connect_timeout,
1105            writer_tx,
1106            state_sink,
1107            controller_lifecycle,
1108            controller_notify,
1109        })
1110    }
1111
1112    /// Returns a cloneable handle to this client's reconnect controller.
1113    #[must_use]
1114    pub fn reconnect_handle(&self) -> SocketReconnectHandle {
1115        SocketReconnectHandle {
1116            connection_mode: Arc::clone(&self.connection_mode),
1117            state_sink: self.state_sink.clone(),
1118            controller_lifecycle: Arc::clone(&self.controller_lifecycle),
1119            controller_notify: Arc::clone(&self.controller_notify),
1120        }
1121    }
1122
1123    /// Requests that the controller replace the active transport.
1124    ///
1125    /// Returns `true` only when this call transitions the client from active to reconnecting.
1126    /// Duplicate, disconnecting, and closed requests return `false`.
1127    #[must_use]
1128    pub fn request_reconnect(&self) -> bool {
1129        self.reconnect_handle().request_reconnect() == ReconnectRequestOutcome::Accepted
1130    }
1131
1132    /// Returns the current connection mode.
1133    #[must_use]
1134    pub fn connection_mode(&self) -> ConnectionMode {
1135        ConnectionMode::from_atomic(&self.connection_mode)
1136    }
1137
1138    /// Returns whether the client connection is active.
1139    ///
1140    /// Returns `true` if the client is connected and has not been signalled to disconnect.
1141    /// The client will automatically retry connection based on its configuration.
1142    #[inline]
1143    #[must_use]
1144    pub fn is_active(&self) -> bool {
1145        self.connection_mode().is_active()
1146    }
1147
1148    /// Returns whether the client is reconnecting.
1149    ///
1150    /// Returns `true` if the client lost connection and is attempting to reestablish it.
1151    /// The client will automatically retry connection based on its configuration.
1152    #[inline]
1153    #[must_use]
1154    pub fn is_reconnecting(&self) -> bool {
1155        self.connection_mode().is_reconnect()
1156    }
1157
1158    /// Returns whether the client is disconnecting.
1159    ///
1160    /// Returns `true` if the client is in disconnect mode.
1161    #[inline]
1162    #[must_use]
1163    pub fn is_disconnecting(&self) -> bool {
1164        self.connection_mode().is_disconnect()
1165    }
1166
1167    /// Returns whether the client is closed.
1168    ///
1169    /// Returns `true` if the client has been explicitly disconnected or reached
1170    /// maximum reconnection attempts. In this state, the client cannot be reused
1171    /// and a new client must be created for further connections.
1172    #[inline]
1173    #[must_use]
1174    pub fn is_closed(&self) -> bool {
1175        self.connection_mode().is_closed()
1176    }
1177
1178    /// Close the client.
1179    ///
1180    /// Controller task will periodically check the disconnect mode
1181    /// and shutdown the client if it is not alive.
1182    pub async fn close(&self) {
1183        self.begin_shutdown();
1184        self.close_with_timeout(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS))
1185            .await;
1186    }
1187
1188    /// Requests shutdown without waiting for the controller task.
1189    pub fn begin_shutdown(&self) {
1190        ConnectionMode::request_disconnect(&self.connection_mode);
1191        self.state_notify.notify_waiters();
1192    }
1193
1194    async fn close_with_timeout(&self, shutdown_timeout: Duration) {
1195        self.begin_shutdown();
1196
1197        if dst::time::timeout(shutdown_timeout, async {
1198            while !self.controller_task.is_finished() {
1199                dst::time::sleep(Duration::from_millis(CONNECTION_STATE_CHECK_INTERVAL_MS)).await;
1200            }
1201        })
1202        .await
1203        .is_err()
1204        {
1205            log::warn!("Timeout waiting for controller task to finish");
1206        }
1207
1208        if !self.controller_task.is_finished() {
1209            self.controller_task.abort();
1210            log_task_aborted("controller");
1211        }
1212
1213        self.connection_mode
1214            .store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1215        self.state_notify.notify_waiters();
1216    }
1217
1218    /// Checks whether the connection is in a terminal state (disconnecting or closed).
1219    ///
1220    /// Single atomic load to fail fast before waiting.
1221    #[inline]
1222    fn check_not_terminal(&self) -> Result<(), SendError> {
1223        match self.connection_mode() {
1224            ConnectionMode::Disconnect | ConnectionMode::Closed => Err(SendError::Closed),
1225            _ => Ok(()),
1226        }
1227    }
1228
1229    /// Waits for the client to become active before sending.
1230    ///
1231    /// Uses `state_notify` for event-driven wakeup so sends resume immediately
1232    /// after reconnection completes. A fallback interval guards against missed
1233    /// notifications.
1234    async fn wait_for_active(&self) -> Result<(), SendError> {
1235        const FALLBACK_INTERVAL_MS: u64 = 100;
1236
1237        let mode = self.connection_mode();
1238        if mode.is_active() {
1239            return Ok(());
1240        }
1241
1242        if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1243            return Err(SendError::Closed);
1244        }
1245
1246        log::debug!("Waiting for client to become ACTIVE before sending...");
1247
1248        let fallback_interval = Duration::from_millis(FALLBACK_INTERVAL_MS);
1249
1250        dst::time::timeout(self.connect_timeout, async {
1251            loop {
1252                // Enable before the state check: an unpolled Notified is unregistered and misses notifies
1253                let mut notified = pin!(self.state_notify.notified());
1254                notified.as_mut().enable();
1255
1256                let mode = self.connection_mode();
1257                if mode.is_active() {
1258                    return Ok(());
1259                }
1260
1261                if matches!(mode, ConnectionMode::Disconnect | ConnectionMode::Closed) {
1262                    return Err(());
1263                }
1264
1265                tokio::select! {
1266                    biased;
1267                    () = notified => {}
1268                    () = dst::time::sleep(fallback_interval) => {}
1269                }
1270            }
1271        })
1272        .await
1273        .map_err(|_| SendError::Timeout)?
1274        .map_err(|()| SendError::Closed)
1275    }
1276
1277    /// Sends a message of the given `data`.
1278    ///
1279    /// Returns `Ok(())` when the message is enqueued to the writer channel. This does not
1280    /// guarantee delivery: if a disconnect occurs concurrently, the writer task may drop the
1281    /// message. During reconnection, messages are buffered and replayed on the new connection.
1282    ///
1283    /// # Errors
1284    ///
1285    /// Returns an error if sending fails.
1286    pub async fn send_bytes(&self, data: Vec<u8>) -> Result<(), SendError> {
1287        self.check_not_terminal()?;
1288        self.wait_for_active().await?;
1289
1290        let msg = WriterCommand::Send(data.into());
1291        self.writer_tx
1292            .send(msg)
1293            .map_err(|e| SendError::BrokenPipe(e.to_string()))
1294    }
1295
1296    fn spawn_controller_task(
1297        mut inner: SocketClientInner,
1298        connection_mode: Arc<AtomicU8>,
1299        state_notify: Arc<tokio::sync::Notify>,
1300        controller_lifecycle: Arc<ControllerLifecycle>,
1301        controller_notify: Arc<tokio::sync::Notify>,
1302        post_reconnection: Option<Arc<dyn Fn() + Send + Sync>>,
1303        reconnect_replay: Option<SocketReconnectReplay>,
1304    ) -> tokio::task::JoinHandle<()> {
1305        const CONTROLLER_FALLBACK_INTERVAL_MS: u64 = 100;
1306
1307        tokio::task::spawn(async move {
1308            let _activity = controller_lifecycle.activity();
1309            log_task_started("controller");
1310
1311            let fallback_interval = Duration::from_millis(CONTROLLER_FALLBACK_INTERVAL_MS);
1312            let mut reconnected_at = None;
1313
1314            loop {
1315                tokio::select! {
1316                    biased;
1317                    () = controller_notify.notified() => {}
1318                    () = state_notify.notified() => {}
1319                    () = dst::time::sleep(fallback_interval) => {}
1320                }
1321
1322                let mut mode = ConnectionMode::from_atomic(&connection_mode);
1323
1324                if mode.is_disconnect() {
1325                    log::debug!("Disconnecting");
1326
1327                    let timeout = Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS);
1328                    if dst::time::timeout(timeout, async {
1329                        // Delay awaiting graceful shutdown
1330                        dst::time::sleep(Duration::from_millis(GRACEFUL_SHUTDOWN_DELAY_MS)).await;
1331
1332                        inner.read_fence.invalidate();
1333                        if !inner.read_task.is_finished() {
1334                            inner.read_task.abort();
1335                            log_task_aborted("read");
1336                        }
1337
1338                        if let Some(task) = &inner.heartbeat_task
1339                            && !task.is_finished()
1340                        {
1341                            task.abort();
1342                            log_task_aborted("heartbeat");
1343                        }
1344                    })
1345                    .await
1346                    .is_err()
1347                    {
1348                        log::warn!("Shutdown timed out after {}s", timeout.as_secs());
1349                    }
1350
1351                    log::debug!("Closed");
1352                    connection_mode.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
1353                    state_notify.notify_waiters();
1354                    break; // Controller finished
1355                }
1356
1357                if mode.is_closed() {
1358                    log::debug!("Connection closed");
1359
1360                    inner.read_fence.invalidate();
1361                    if !inner.read_task.is_finished() {
1362                        inner.read_task.abort();
1363                        log_task_aborted("read");
1364                    }
1365
1366                    if let Some(task) = &inner.heartbeat_task
1367                        && !task.is_finished()
1368                    {
1369                        task.abort();
1370                        log_task_aborted("heartbeat");
1371                    }
1372
1373                    state_notify.notify_waiters();
1374                    break;
1375                }
1376
1377                if mode.is_active() && !inner.is_alive() {
1378                    if ConnectionMode::request_reconnect_with_sink(
1379                        &connection_mode,
1380                        inner.state_sink.as_ref(),
1381                    ) {
1382                        log::info!("Detected dead read task, transitioning to RECONNECT");
1383                    }
1384                    mode = ConnectionMode::from_atomic(&connection_mode);
1385                }
1386
1387                if mode.is_reconnect() {
1388                    let reconnect_uptime = reconnected_at
1389                        .take()
1390                        .map(|started: dst::time::Instant| started.elapsed());
1391                    let previous_reconnect_stable = reconnect_uptime
1392                        .is_some_and(|uptime| uptime >= RECONNECT_STABILITY_THRESHOLD);
1393
1394                    if previous_reconnect_stable {
1395                        inner.backoff.reset();
1396                        inner.reconnect_attempt_count = 0;
1397                        log::debug!(
1398                            "Socket remained active for at least {}s, resetting reconnect cycle",
1399                            RECONNECT_STABILITY_THRESHOLD.as_secs()
1400                        );
1401                    }
1402
1403                    // Check max reconnection attempts before attempting reconnect
1404                    if let Some(max_attempts) = inner.reconnect_max_attempts
1405                        && inner.reconnect_attempt_count >= max_attempts
1406                    {
1407                        log::error!(
1408                            "Max reconnection attempts ({max_attempts}) exceeded, transitioning to CLOSED"
1409                        );
1410
1411                        if connection_mode
1412                            .compare_exchange(
1413                                ConnectionMode::Reconnect.as_u8(),
1414                                ConnectionMode::Closed.as_u8(),
1415                                Ordering::SeqCst,
1416                                Ordering::SeqCst,
1417                            )
1418                            .is_ok()
1419                        {
1420                            state_notify.notify_waiters();
1421                            break;
1422                        }
1423                        continue;
1424                    }
1425
1426                    let backoff_delay = if reconnect_uptime.is_some() && !previous_reconnect_stable
1427                    {
1428                        inner.backoff.next_duration()
1429                    } else {
1430                        Duration::ZERO
1431                    };
1432
1433                    let duration = inner.reconnect_throttle.gated_delay(backoff_delay);
1434                    if !duration.is_zero() {
1435                        log::warn!("Backing off for {}s...", duration.as_secs_f64());
1436
1437                        if !wait_reconnect_delay(
1438                            duration,
1439                            connection_mode.as_ref(),
1440                            state_notify.as_ref(),
1441                        )
1442                        .await
1443                        {
1444                            log::debug!("Backoff interrupted by terminal state");
1445                            continue;
1446                        }
1447                    }
1448
1449                    inner.reconnect_attempt_count += 1;
1450                    inner.reconnect_throttle.record_attempt();
1451
1452                    // Race reconnect against disconnect notification
1453                    let reconnect_result = tokio::select! {
1454                        biased;
1455                        result = inner.reconnect(reconnect_replay.as_ref()) => Some(result),
1456                        () = async {
1457                            loop {
1458                                // Enable before the check so a disconnect notify between iterations is not missed
1459                                let mut notified = pin!(state_notify.notified());
1460                                notified.as_mut().enable();
1461
1462                                if ConnectionMode::from_atomic(&connection_mode).is_disconnect() {
1463                                    break;
1464                                }
1465                                notified.await;
1466                            }
1467                        } => None,
1468                    };
1469
1470                    match reconnect_result {
1471                        None => {
1472                            log::debug!("Reconnect interrupted by disconnect");
1473                        }
1474                        Some(Ok(ReconnectOutcome::Reconnected)) => {
1475                            log::debug!("Reconnected successfully");
1476                            reconnected_at = Some(dst::time::Instant::now());
1477
1478                            state_notify.notify_waiters();
1479
1480                            // The outcome records a completed reconnection; emit recovery
1481                            // callbacks only while the replacement is still `Active`, not
1482                            // after a teardown or another drop.
1483                            if ConnectionMode::from_atomic(&connection_mode).is_active() {
1484                                if let Some(ref handler) = post_reconnection {
1485                                    handler();
1486                                    log::debug!("Called `post_reconnection` handler");
1487                                }
1488                            } else {
1489                                log::debug!(
1490                                    "Skipping post_reconnection handlers due to disconnect state"
1491                                );
1492                            }
1493                        }
1494                        Some(Ok(ReconnectOutcome::Aborted)) => {
1495                            log::debug!("Reconnect aborted");
1496                        }
1497                        Some(Err(e)) => {
1498                            let duration = inner.backoff.next_duration();
1499                            log::warn!(
1500                                "Reconnect attempt {} failed: {e}",
1501                                inner.reconnect_attempt_count
1502                            );
1503
1504                            if !duration.is_zero() {
1505                                log::warn!("Backing off for {}s...", duration.as_secs_f64());
1506                                if !wait_reconnect_delay(
1507                                    duration,
1508                                    connection_mode.as_ref(),
1509                                    state_notify.as_ref(),
1510                                )
1511                                .await
1512                                {
1513                                    log::debug!("Backoff interrupted by terminal state");
1514                                }
1515                            }
1516                        }
1517                    }
1518                }
1519            }
1520            log_task_stopped("controller");
1521        })
1522    }
1523}
1524
1525// Dropping cancels background work without reporting a terminal socket state transition.
1526impl Drop for SocketClient {
1527    fn drop(&mut self) {
1528        let controller_running = !self.controller_task.is_finished();
1529        self.controller_lifecycle.close_and_abort();
1530
1531        if controller_running {
1532            log_task_aborted("controller");
1533        }
1534    }
1535}
1536
1537#[cfg(test)]
1538#[cfg(not(feature = "turmoil"))]
1539#[cfg(not(all(feature = "simulation", madsim)))] // transport-layer I/O not simulated
1540#[cfg(target_os = "linux")] // Only run network tests on Linux (CI stability)
1541mod tests {
1542    use nautilus_common::testing::wait_until_async;
1543    use parking_lot::Mutex as BlockingMutex;
1544    use rstest::rstest;
1545    use tokio::{
1546        io::{AsyncReadExt, AsyncWriteExt},
1547        net::{TcpListener, TcpStream},
1548        sync::Mutex,
1549        task,
1550        time::{Duration, sleep},
1551    };
1552
1553    use super::*;
1554    use crate::{SocketState, socket::SocketHeartbeat};
1555
1556    async fn bind_test_server() -> (u16, TcpListener) {
1557        let listener = TcpListener::bind("127.0.0.1:0")
1558            .await
1559            .expect("Failed to bind ephemeral port");
1560        let port = listener.local_addr().unwrap().port();
1561        (port, listener)
1562    }
1563
1564    async fn run_echo_server(mut socket: TcpStream) {
1565        let mut buf = Vec::new();
1566        loop {
1567            match socket.read_buf(&mut buf).await {
1568                Ok(0) => {
1569                    break;
1570                }
1571                Ok(_n) => {
1572                    while let Some(idx) = buf.array_windows().position(|w| w == b"\r\n") {
1573                        let mut line = buf.drain(..idx + 2).collect::<Vec<u8>>();
1574                        // Remove trailing \r\n
1575                        line.truncate(line.len() - 2);
1576
1577                        if line == b"close" {
1578                            let _ = socket.shutdown().await;
1579                            return;
1580                        }
1581
1582                        let mut echo_data = line;
1583                        echo_data.extend_from_slice(b"\r\n");
1584                        if socket.write_all(&echo_data).await.is_err() {
1585                            break;
1586                        }
1587                    }
1588                }
1589                Err(e) => {
1590                    eprintln!("Server read error: {e}");
1591                    break;
1592                }
1593            }
1594        }
1595    }
1596
1597    #[tokio::test]
1598    async fn test_basic_send_receive() {
1599        let (port, listener) = bind_test_server().await;
1600        let server_task = task::spawn(async move {
1601            let (socket, _) = listener.accept().await.unwrap();
1602            run_echo_server(socket).await;
1603        });
1604
1605        let config = SocketConfig {
1606            url: format!("127.0.0.1:{port}"),
1607            mode: Mode::Plain,
1608            suffix: b"\r\n".to_vec(),
1609            message_handler: None,
1610            heartbeat: None,
1611            connect_timeout_ms: None,
1612            reconnect_delay_initial_ms: None,
1613            reconnect_backoff_factor: None,
1614            reconnect_delay_max_ms: None,
1615            reconnect_jitter_ms: None,
1616            reconnect_max_attempts: None,
1617            connection_max_retries: None,
1618            heartbeat_timeout_secs: None,
1619            certs_dir: None,
1620        };
1621
1622        let client = SocketClient::builder()
1623            .config(config)
1624            .connect()
1625            .await
1626            .expect("Client connect failed unexpectedly");
1627
1628        client.send_bytes(b"Hello".into()).await.unwrap();
1629        client.send_bytes(b"World".into()).await.unwrap();
1630
1631        // Wait a bit for the server to echo them back
1632        sleep(Duration::from_millis(100)).await;
1633
1634        client.send_bytes(b"close".into()).await.unwrap();
1635        server_task.await.unwrap();
1636        assert!(!client.is_closed());
1637    }
1638
1639    #[tokio::test]
1640    async fn test_reconnect_fail_exhausted() {
1641        let (port, listener) = bind_test_server().await;
1642        drop(listener); // We drop it immediately -> no server is listening
1643
1644        // Wait until port is truly unavailable (OS has released it)
1645        wait_until_async(
1646            || async {
1647                TcpStream::connect(format!("127.0.0.1:{port}"))
1648                    .await
1649                    .is_err()
1650            },
1651            Duration::from_secs(2),
1652        )
1653        .await;
1654
1655        let config = SocketConfig {
1656            url: format!("127.0.0.1:{port}"),
1657            mode: Mode::Plain,
1658            suffix: b"\r\n".to_vec(),
1659            message_handler: None,
1660            heartbeat: None,
1661            connect_timeout_ms: Some(100),
1662            reconnect_delay_initial_ms: Some(50),
1663            reconnect_backoff_factor: Some(1.0),
1664            reconnect_delay_max_ms: Some(50),
1665            reconnect_jitter_ms: Some(0),
1666            connection_max_retries: Some(1),
1667            reconnect_max_attempts: None,
1668            heartbeat_timeout_secs: None,
1669            certs_dir: None,
1670        };
1671
1672        let client_res = SocketClient::builder().config(config).connect().await;
1673        assert!(
1674            client_res.is_err(),
1675            "Should fail quickly with no server listening"
1676        );
1677    }
1678
1679    #[tokio::test]
1680    async fn test_user_disconnect() {
1681        let (port, listener) = bind_test_server().await;
1682        let server_task = task::spawn(async move {
1683            let (socket, _) = listener.accept().await.unwrap();
1684            let mut buf = [0u8; 1024];
1685            let _ = socket.try_read(&mut buf);
1686
1687            loop {
1688                sleep(Duration::from_secs(1)).await;
1689            }
1690        });
1691
1692        let config = SocketConfig {
1693            url: format!("127.0.0.1:{port}"),
1694            mode: Mode::Plain,
1695            suffix: b"\r\n".to_vec(),
1696            message_handler: None,
1697            heartbeat: None,
1698            connect_timeout_ms: None,
1699            reconnect_delay_initial_ms: None,
1700            reconnect_backoff_factor: None,
1701            reconnect_delay_max_ms: None,
1702            reconnect_jitter_ms: None,
1703            reconnect_max_attempts: None,
1704            connection_max_retries: None,
1705            heartbeat_timeout_secs: None,
1706            certs_dir: None,
1707        };
1708
1709        let client = SocketClient::builder()
1710            .config(config)
1711            .connect()
1712            .await
1713            .unwrap();
1714
1715        client.close().await;
1716        assert!(client.is_closed());
1717        server_task.abort();
1718    }
1719
1720    #[tokio::test]
1721    async fn test_close_after_closed_returns_fast_and_preserves_state() {
1722        let (port, listener) = bind_test_server().await;
1723
1724        let server_task = task::spawn(async move {
1725            // Accept the first connection then drop it; never accept again so
1726            // the client exhausts its reconnect attempts and transitions to CLOSED
1727            let (socket, _) = listener.accept().await.unwrap();
1728            drop(socket);
1729            drop(listener);
1730            sleep(Duration::from_secs(5)).await;
1731        });
1732
1733        let config = SocketConfig {
1734            url: format!("127.0.0.1:{port}"),
1735            mode: Mode::Plain,
1736            suffix: b"\r\n".to_vec(),
1737            message_handler: None,
1738            heartbeat: None,
1739            connect_timeout_ms: Some(200),
1740            reconnect_delay_initial_ms: Some(50),
1741            reconnect_backoff_factor: Some(1.0),
1742            reconnect_delay_max_ms: Some(50),
1743            reconnect_jitter_ms: Some(0),
1744            connection_max_retries: None,
1745            reconnect_max_attempts: Some(1),
1746            heartbeat_timeout_secs: None,
1747            certs_dir: None,
1748        };
1749
1750        let client = SocketClient::builder()
1751            .config(config)
1752            .connect()
1753            .await
1754            .unwrap();
1755
1756        wait_until_async(|| async { client.is_closed() }, Duration::from_secs(5)).await;
1757
1758        // Closing an already CLOSED client must return promptly (no 5s spin
1759        // waiting for a controller that has already exited) and must not
1760        // regress the terminal state to DISCONNECT
1761        let start = std::time::Instant::now();
1762        client.close().await;
1763        let elapsed = start.elapsed();
1764
1765        assert!(client.is_closed(), "Client should remain CLOSED");
1766        assert!(
1767            !client.is_disconnecting(),
1768            "Closed client should not report DISCONNECT after close()"
1769        );
1770        assert!(
1771            elapsed < Duration::from_secs(2),
1772            "close() on a closed client should return fast, took {elapsed:?}"
1773        );
1774
1775        server_task.abort();
1776    }
1777
1778    #[tokio::test]
1779    async fn test_heartbeat() {
1780        let (port, listener) = bind_test_server().await;
1781        let received = Arc::new(Mutex::new(Vec::new()));
1782        let received2 = received.clone();
1783
1784        let server_task = task::spawn(async move {
1785            let (socket, _) = listener.accept().await.unwrap();
1786
1787            let mut buf = Vec::new();
1788            loop {
1789                match socket.try_read_buf(&mut buf) {
1790                    Ok(0) => break,
1791                    Ok(_) => {
1792                        while let Some(idx) = buf.array_windows().position(|w| w == b"\r\n") {
1793                            let mut line = buf.drain(..idx + 2).collect::<Vec<u8>>();
1794                            line.truncate(line.len() - 2);
1795                            received2.lock().await.push(line);
1796                        }
1797                    }
1798                    Err(_) => {
1799                        tokio::time::sleep(Duration::from_millis(10)).await;
1800                    }
1801                }
1802            }
1803        });
1804
1805        let config = SocketConfig {
1806            url: format!("127.0.0.1:{port}"),
1807            mode: Mode::Plain,
1808            suffix: b"\r\n".to_vec(),
1809            message_handler: None,
1810            heartbeat: Some(SocketHeartbeat {
1811                interval_secs: 1,
1812                payload: b"ping".to_vec(),
1813            }),
1814            connect_timeout_ms: None,
1815            reconnect_delay_initial_ms: None,
1816            reconnect_backoff_factor: None,
1817            reconnect_delay_max_ms: None,
1818            reconnect_jitter_ms: None,
1819            reconnect_max_attempts: None,
1820            connection_max_retries: None,
1821            heartbeat_timeout_secs: None,
1822            certs_dir: None,
1823        };
1824
1825        let client = SocketClient::builder()
1826            .config(config)
1827            .connect()
1828            .await
1829            .unwrap();
1830
1831        // Wait ~3 seconds to collect some heartbeats
1832        sleep(Duration::from_secs(3)).await;
1833
1834        {
1835            let lock = received.lock().await;
1836            let pings = lock
1837                .iter()
1838                .filter(|line| line == &&b"ping".to_vec())
1839                .count();
1840            assert!(
1841                pings >= 2,
1842                "Expected at least 2 heartbeat pings; got {pings}"
1843            );
1844        }
1845
1846        client.close().await;
1847        server_task.abort();
1848    }
1849
1850    #[tokio::test]
1851    async fn test_reconnect_success() {
1852        let (port, listener) = bind_test_server().await;
1853
1854        // Spawn a server task that:
1855        // 1. Accepts the first connection and then drops it after a short delay (simulate disconnect)
1856        // 2. Waits a bit and then accepts a new connection and runs the echo server
1857        let server_task = task::spawn(async move {
1858            // Accept first connection
1859            let (mut socket, _) = listener.accept().await.expect("First accept failed");
1860
1861            // Wait briefly and then force-close the connection
1862            sleep(Duration::from_millis(500)).await;
1863            let _ = socket.shutdown().await;
1864
1865            // Wait for the client's reconnect attempt
1866            sleep(Duration::from_millis(500)).await;
1867
1868            // Run the echo server on the new connection
1869            let (socket, _) = listener.accept().await.expect("Second accept failed");
1870            run_echo_server(socket).await;
1871        });
1872
1873        let config = SocketConfig {
1874            url: format!("127.0.0.1:{port}"),
1875            mode: Mode::Plain,
1876            suffix: b"\r\n".to_vec(),
1877            message_handler: None,
1878            heartbeat: None,
1879            connect_timeout_ms: Some(5_000),
1880            reconnect_delay_initial_ms: Some(500),
1881            reconnect_delay_max_ms: Some(5_000),
1882            reconnect_backoff_factor: Some(2.0),
1883            reconnect_jitter_ms: Some(50),
1884            reconnect_max_attempts: None,
1885            connection_max_retries: None,
1886            heartbeat_timeout_secs: None,
1887            certs_dir: None,
1888        };
1889
1890        let client = SocketClient::builder()
1891            .config(config)
1892            .connect()
1893            .await
1894            .expect("Client connect failed unexpectedly");
1895
1896        // Initially, the client should be active
1897        assert!(client.is_active(), "Client should start as active");
1898
1899        // Wait until the client loses connection (i.e. not active),
1900        // then wait until it reconnects (active again).
1901        wait_until_async(|| async { client.is_active() }, Duration::from_secs(10)).await;
1902
1903        client
1904            .send_bytes(b"TestReconnect".into())
1905            .await
1906            .expect("Send failed");
1907
1908        client.close().await;
1909        server_task.abort();
1910    }
1911
1912    #[rstest]
1913    #[tokio::test]
1914    async fn test_state_sink_reports_initial_loss_and_recovery() {
1915        let (port, listener) = bind_test_server().await;
1916        let server_task = task::spawn(async move {
1917            let (mut socket, _) = listener.accept().await.unwrap();
1918            sleep(Duration::from_millis(100)).await;
1919            socket.shutdown().await.unwrap();
1920
1921            let (socket, _) = listener.accept().await.unwrap();
1922            run_echo_server(socket).await;
1923        });
1924        let config = SocketConfig {
1925            url: format!("127.0.0.1:{port}"),
1926            mode: Mode::Plain,
1927            suffix: b"\r\n".to_vec(),
1928            message_handler: None,
1929            heartbeat: None,
1930            connect_timeout_ms: Some(1_000),
1931            reconnect_delay_initial_ms: Some(10),
1932            reconnect_backoff_factor: Some(1.0),
1933            reconnect_delay_max_ms: Some(10),
1934            reconnect_jitter_ms: Some(0),
1935            reconnect_max_attempts: Some(3),
1936            connection_max_retries: Some(1),
1937            heartbeat_timeout_secs: None,
1938            certs_dir: None,
1939        };
1940        let states = Arc::new(BlockingMutex::new(Vec::new()));
1941        let states_callback = Arc::clone(&states);
1942        let sink = SocketStateSink::new(move |state| {
1943            states_callback.lock().push(state);
1944        });
1945
1946        let client = SocketClient::builder()
1947            .config(config)
1948            .state_sink(sink)
1949            .connect()
1950            .await
1951            .unwrap();
1952
1953        assert_eq!(*states.lock(), vec![SocketState::Connected]);
1954
1955        wait_until_async(
1956            || {
1957                let states = Arc::clone(&states);
1958                async move { states.lock().len() == 3 }
1959            },
1960            Duration::from_secs(5),
1961        )
1962        .await;
1963        assert_eq!(
1964            *states.lock(),
1965            vec![
1966                SocketState::Connected,
1967                SocketState::Disconnected,
1968                SocketState::Connected,
1969            ]
1970        );
1971
1972        client.close().await;
1973        assert_eq!(states.lock().len(), 3);
1974        server_task.abort();
1975    }
1976
1977    #[rstest]
1978    #[tokio::test]
1979    async fn test_state_sink_ignores_initial_connection_failure() {
1980        let (port, listener) = bind_test_server().await;
1981        drop(listener);
1982        let config = SocketConfig {
1983            url: format!("127.0.0.1:{port}"),
1984            mode: Mode::Plain,
1985            suffix: b"\r\n".to_vec(),
1986            message_handler: None,
1987            heartbeat: None,
1988            connect_timeout_ms: Some(100),
1989            reconnect_delay_initial_ms: Some(1),
1990            reconnect_backoff_factor: Some(1.0),
1991            reconnect_delay_max_ms: Some(1),
1992            reconnect_jitter_ms: Some(0),
1993            reconnect_max_attempts: Some(1),
1994            connection_max_retries: Some(1),
1995            heartbeat_timeout_secs: None,
1996            certs_dir: None,
1997        };
1998        let states = Arc::new(BlockingMutex::new(Vec::new()));
1999        let states_callback = Arc::clone(&states);
2000        let sink = SocketStateSink::new(move |state| {
2001            states_callback.lock().push(state);
2002        });
2003
2004        let result = SocketClient::builder()
2005            .config(config)
2006            .state_sink(sink)
2007            .connect()
2008            .await;
2009
2010        assert!(result.is_err());
2011        assert_eq!(*states.lock(), Vec::new());
2012    }
2013}
2014
2015#[cfg(test)]
2016#[cfg(not(feature = "turmoil"))]
2017#[cfg(not(all(feature = "simulation", madsim)))] // transport-layer I/O not simulated
2018mod rust_tests {
2019    use std::{
2020        pin::Pin,
2021        sync::{
2022            Arc,
2023            atomic::{AtomicBool, AtomicUsize},
2024        },
2025        task::{Context, Poll, Waker},
2026    };
2027
2028    use nautilus_common::testing::wait_until_async;
2029    use parking_lot::{Condvar, Mutex};
2030    use rstest::rstest;
2031    use tokio::{
2032        io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream, ReadBuf},
2033        net::TcpListener,
2034        sync::oneshot,
2035        task::{self, yield_now},
2036        time::{Duration, sleep},
2037    };
2038
2039    use super::*;
2040    use crate::{SocketState, socket::SocketHeartbeat};
2041
2042    const TEST_TIMEOUT: Duration = Duration::from_secs(10);
2043
2044    struct CondvarReleaseGuard {
2045        release: Arc<(Mutex<bool>, Condvar)>,
2046    }
2047
2048    impl CondvarReleaseGuard {
2049        fn new(release: Arc<(Mutex<bool>, Condvar)>) -> Self {
2050            Self { release }
2051        }
2052
2053        fn release(&self) {
2054            let (lock, condvar) = self.release.as_ref();
2055            let mut released = lock.lock();
2056            *released = true;
2057            condvar.notify_all();
2058        }
2059    }
2060
2061    impl Drop for CondvarReleaseGuard {
2062        fn drop(&mut self) {
2063            self.release();
2064        }
2065    }
2066
2067    async fn recv_rendezvous<T: Send + 'static>(
2068        receiver: std::sync::mpsc::Receiver<T>,
2069        name: &'static str,
2070    ) -> T {
2071        let receive_task = tokio::task::spawn_blocking(move || receiver.recv_timeout(TEST_TIMEOUT));
2072
2073        match tokio::time::timeout(TEST_TIMEOUT * 2, receive_task).await {
2074            Ok(Ok(Ok(value))) => value,
2075            Ok(Ok(Err(e))) => {
2076                panic!("{name} did not arrive within the test timeout: {e}")
2077            }
2078            Ok(Err(e)) => panic!("{name} receive task failed: {e}"),
2079            Err(e) => panic!("{name} receive task did not finish: {e}"),
2080        }
2081    }
2082
2083    async fn await_task_termination(
2084        task: tokio::task::JoinHandle<()>,
2085        name: &'static str,
2086        cancellation_ok: bool,
2087    ) {
2088        match tokio::time::timeout(TEST_TIMEOUT, task).await {
2089            Ok(Ok(())) => {}
2090            Ok(Err(e)) if cancellation_ok && e.is_cancelled() => {}
2091            Ok(Err(e)) => panic!("{name} failed: {e}"),
2092            Err(e) => panic!("{name} did not terminate within the test timeout: {e}"),
2093        }
2094    }
2095
2096    fn reconnect_test_config(port: u16) -> SocketConfig {
2097        SocketConfig {
2098            url: format!("127.0.0.1:{port}"),
2099            mode: Mode::Plain,
2100            suffix: b"\r\n".to_vec(),
2101            message_handler: None,
2102            heartbeat: None,
2103            connect_timeout_ms: Some(1_000),
2104            reconnect_delay_initial_ms: None,
2105            reconnect_backoff_factor: None,
2106            reconnect_delay_max_ms: None,
2107            reconnect_jitter_ms: None,
2108            connection_max_retries: Some(1),
2109            reconnect_max_attempts: None,
2110            heartbeat_timeout_secs: None,
2111            certs_dir: None,
2112        }
2113    }
2114
2115    #[rstest]
2116    #[tokio::test]
2117    async fn test_reconnect_outcome_is_aborted_before_connect() {
2118        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2119        let port = listener.local_addr().unwrap().port();
2120        let server = tokio::spawn(async move {
2121            let (_socket, _) = listener.accept().await.unwrap();
2122            std::future::pending::<()>().await;
2123        });
2124        let mut inner = SocketClientInner::connect_url(reconnect_test_config(port), None)
2125            .await
2126            .unwrap();
2127        inner
2128            .connection_mode
2129            .store(ConnectionMode::Disconnect.as_u8(), Ordering::SeqCst);
2130
2131        let outcome = inner.reconnect(None).await.unwrap();
2132
2133        assert_eq!(outcome, ReconnectOutcome::Aborted);
2134        server.abort();
2135    }
2136
2137    #[rstest]
2138    #[tokio::test]
2139    async fn test_reconnect_outcome_is_reconnected() {
2140        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2141        let port = listener.local_addr().unwrap().port();
2142        let server = tokio::spawn(async move {
2143            let (_first, _) = listener.accept().await.unwrap();
2144            let (_second, _) = listener.accept().await.unwrap();
2145            std::future::pending::<()>().await;
2146        });
2147        let mut inner = SocketClientInner::connect_url(reconnect_test_config(port), None)
2148            .await
2149            .unwrap();
2150        inner
2151            .connection_mode
2152            .store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2153
2154        let outcome = inner.reconnect(None).await.unwrap();
2155
2156        assert_eq!(outcome, ReconnectOutcome::Reconnected);
2157        assert_eq!(
2158            ConnectionMode::from_atomic(&inner.connection_mode),
2159            ConnectionMode::Active
2160        );
2161        server.abort();
2162    }
2163
2164    #[rstest]
2165    #[tokio::test]
2166    async fn test_inner_drop_invalidates_read_fence_and_aborts_tasks() {
2167        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2168        let port = listener.local_addr().unwrap().port();
2169        let server = tokio::spawn(async move {
2170            let (_socket, _) = listener.accept().await.unwrap();
2171            std::future::pending::<()>().await;
2172        });
2173        let mut config = reconnect_test_config(port);
2174        config.heartbeat = Some(SocketHeartbeat {
2175            interval_secs: 60,
2176            payload: b"ping\r\n".to_vec(),
2177        });
2178        let inner = SocketClientInner::connect_url(config, None).await.unwrap();
2179        let read_fence = inner.read_fence.clone();
2180        let read_abort = inner.read_task.abort_handle();
2181        let write_abort = inner.write_task.abort_handle();
2182        let heartbeat_abort = inner
2183            .heartbeat_task
2184            .as_ref()
2185            .expect("heartbeat task should be spawned for a configured heartbeat")
2186            .abort_handle();
2187
2188        assert!(read_fence.is_valid(), "read fence should start valid");
2189        assert!(
2190            !read_abort.is_finished(),
2191            "read task should be running before drop"
2192        );
2193        assert!(
2194            !write_abort.is_finished(),
2195            "write task should be running before drop"
2196        );
2197        assert!(
2198            !heartbeat_abort.is_finished(),
2199            "heartbeat task should be running before drop"
2200        );
2201
2202        drop(inner);
2203        wait_until_async(
2204            || async {
2205                read_abort.is_finished()
2206                    && write_abort.is_finished()
2207                    && heartbeat_abort.is_finished()
2208            },
2209            TEST_TIMEOUT,
2210        )
2211        .await;
2212
2213        assert!(!read_fence.is_valid(), "read fence was not invalidated");
2214        assert!(read_abort.is_finished(), "read task was not aborted");
2215        assert!(write_abort.is_finished(), "write task was not aborted");
2216        assert!(
2217            heartbeat_abort.is_finished(),
2218            "heartbeat task was not aborted"
2219        );
2220        server.abort();
2221    }
2222
2223    struct ScriptedReader {
2224        first: Option<Vec<u8>>,
2225        remainder: Arc<Mutex<Option<Vec<u8>>>>,
2226        pending_tx: Option<oneshot::Sender<()>>,
2227        waker: Arc<Mutex<Option<Waker>>>,
2228    }
2229
2230    impl AsyncRead for ScriptedReader {
2231        fn poll_read(
2232            mut self: Pin<&mut Self>,
2233            cx: &mut Context<'_>,
2234            buf: &mut ReadBuf<'_>,
2235        ) -> Poll<std::io::Result<()>> {
2236            if let Some(first) = self.first.take() {
2237                buf.put_slice(&first);
2238                return Poll::Ready(Ok(()));
2239            }
2240
2241            if let Some(remainder) = self.remainder.lock().take() {
2242                buf.put_slice(&remainder);
2243                return Poll::Ready(Ok(()));
2244            }
2245
2246            if let Some(tx) = self.pending_tx.take() {
2247                let _ = tx.send(());
2248            }
2249            *self.waker.lock() = Some(cx.waker().clone());
2250            Poll::Pending
2251        }
2252    }
2253
2254    struct LimitProbeReader {
2255        remaining: usize,
2256        read_after_limit: Arc<AtomicBool>,
2257    }
2258
2259    impl AsyncRead for LimitProbeReader {
2260        fn poll_read(
2261            mut self: Pin<&mut Self>,
2262            _cx: &mut Context<'_>,
2263            buf: &mut ReadBuf<'_>,
2264        ) -> Poll<std::io::Result<()>> {
2265            if self.remaining == 0 {
2266                self.read_after_limit.store(true, Ordering::SeqCst);
2267                return Poll::Ready(Ok(()));
2268            }
2269
2270            let read_len = self.remaining.min(buf.remaining());
2271            buf.initialize_unfilled_to(read_len).fill(b'x');
2272            buf.advance(read_len);
2273            self.remaining -= read_len;
2274            Poll::Ready(Ok(()))
2275        }
2276    }
2277
2278    struct BackpressuredWriter {
2279        stream: DuplexStream,
2280        pending_tx: Option<oneshot::Sender<()>>,
2281    }
2282
2283    impl AsyncWrite for BackpressuredWriter {
2284        fn poll_write(
2285            mut self: Pin<&mut Self>,
2286            cx: &mut Context<'_>,
2287            buf: &[u8],
2288        ) -> Poll<std::io::Result<usize>> {
2289            let result = Pin::new(&mut self.stream).poll_write(cx, buf);
2290            if result.is_pending()
2291                && let Some(tx) = self.pending_tx.take()
2292            {
2293                let _ = tx.send(());
2294            }
2295            result
2296        }
2297
2298        fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2299            Pin::new(&mut self.stream).poll_flush(cx)
2300        }
2301
2302        fn poll_shutdown(
2303            mut self: Pin<&mut Self>,
2304            cx: &mut Context<'_>,
2305        ) -> Poll<std::io::Result<()>> {
2306            Pin::new(&mut self.stream).poll_shutdown(cx)
2307        }
2308    }
2309
2310    struct FailingWriter;
2311
2312    impl AsyncWrite for FailingWriter {
2313        fn poll_write(
2314            self: Pin<&mut Self>,
2315            _cx: &mut Context<'_>,
2316            _buf: &[u8],
2317        ) -> Poll<std::io::Result<usize>> {
2318            Poll::Ready(Err(std::io::Error::new(
2319                std::io::ErrorKind::BrokenPipe,
2320                "test writer failure",
2321            )))
2322        }
2323
2324        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2325            Poll::Ready(Ok(()))
2326        }
2327
2328        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2329            Poll::Ready(Ok(()))
2330        }
2331    }
2332
2333    struct RecordingWriter {
2334        bytes: Arc<Mutex<Vec<u8>>>,
2335    }
2336
2337    impl AsyncWrite for RecordingWriter {
2338        fn poll_write(
2339            self: Pin<&mut Self>,
2340            _cx: &mut Context<'_>,
2341            buf: &[u8],
2342        ) -> Poll<std::io::Result<usize>> {
2343            self.bytes.lock().extend_from_slice(buf);
2344            Poll::Ready(Ok(buf.len()))
2345        }
2346
2347        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2348            Poll::Ready(Ok(()))
2349        }
2350
2351        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2352            Poll::Ready(Ok(()))
2353        }
2354    }
2355
2356    fn test_socket_client(
2357        connection_state: Arc<AtomicU8>,
2358        state_notify: Arc<tokio::sync::Notify>,
2359        controller_task: tokio::task::JoinHandle<()>,
2360    ) -> SocketClient {
2361        let (writer_tx, _writer_rx) = tokio::sync::mpsc::unbounded_channel();
2362        let controller_lifecycle = Arc::new(ControllerLifecycle::new());
2363        controller_lifecycle.set_abort_handle(controller_task.abort_handle());
2364
2365        SocketClient {
2366            controller_task,
2367            connection_mode: connection_state,
2368            state_notify,
2369            connect_timeout: Duration::from_secs(1),
2370            writer_tx,
2371            controller_lifecycle,
2372            controller_notify: Arc::new(tokio::sync::Notify::new()),
2373            state_sink: None,
2374        }
2375    }
2376
2377    #[rstest]
2378    #[tokio::test]
2379    async fn test_reconnect_handle_is_closed_after_client_drop() {
2380        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2381        let state_notify = Arc::new(tokio::sync::Notify::new());
2382        let controller_task = tokio::spawn(std::future::pending::<()>());
2383        let client = test_socket_client(connection_state, state_notify, controller_task);
2384        let handle = client.reconnect_handle();
2385
2386        drop(client);
2387
2388        assert_eq!(handle.request_reconnect(), ReconnectRequestOutcome::Closed);
2389    }
2390
2391    #[rstest]
2392    #[tokio::test]
2393    async fn test_concurrent_drop_defers_controller_abort_until_request_completes() {
2394        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2395        let state_notify = Arc::new(tokio::sync::Notify::new());
2396
2397        let controller_task = tokio::spawn(std::future::pending::<()>());
2398        let controller_abort = controller_task.abort_handle();
2399        let mut client = test_socket_client(connection_state, state_notify, controller_task);
2400        let callback_release = Arc::new((Mutex::new(false), Condvar::new()));
2401        let callback_release_guard = CondvarReleaseGuard::new(Arc::clone(&callback_release));
2402        let callback_release_clone = Arc::clone(&callback_release);
2403        let (callback_entered_tx, callback_entered_rx) = std::sync::mpsc::channel();
2404        let states = Arc::new(Mutex::new(Vec::new()));
2405        let states_clone = Arc::clone(&states);
2406        client.state_sink = Some(SocketStateSink::new(move |state| {
2407            states_clone.lock().push(state);
2408            callback_entered_tx.send(()).unwrap();
2409            let (lock, condvar) = callback_release_clone.as_ref();
2410            let mut released = lock.lock();
2411
2412            while !*released {
2413                condvar.wait(&mut released);
2414            }
2415        }));
2416        let handle = client.reconnect_handle();
2417        let surviving_handle = handle.clone();
2418        let controller_notify = Arc::clone(&handle.controller_notify);
2419        let request_thread = std::thread::spawn(move || handle.request_reconnect());
2420
2421        recv_rendezvous(callback_entered_rx, "socket state callback entry").await;
2422        let (drop_finished_tx, drop_finished_rx) = std::sync::mpsc::channel();
2423        let drop_thread = std::thread::spawn(move || {
2424            drop(client);
2425            drop_finished_tx.send(()).unwrap();
2426        });
2427
2428        recv_rendezvous(drop_finished_rx, "client drop").await;
2429        drop_thread.join().unwrap();
2430        assert!(!controller_abort.is_finished());
2431
2432        callback_release_guard.release();
2433        assert_eq!(
2434            request_thread.join().unwrap(),
2435            ReconnectRequestOutcome::Accepted
2436        );
2437        tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2438            .await
2439            .expect("accepted request should notify before deferred controller abort");
2440        wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
2441
2442        assert_eq!(
2443            surviving_handle.request_reconnect(),
2444            ReconnectRequestOutcome::Closed
2445        );
2446        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2447        assert!(
2448            tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2449                .await
2450                .is_err(),
2451            "closed request should not notify controller",
2452        );
2453    }
2454
2455    #[rstest]
2456    #[tokio::test]
2457    async fn test_reconnect_state_callback_can_drop_client() {
2458        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2459        let state_notify = Arc::new(tokio::sync::Notify::new());
2460
2461        let controller_task = tokio::spawn(std::future::pending::<()>());
2462        let controller_abort = controller_task.abort_handle();
2463        let mut client = test_socket_client(connection_state, state_notify, controller_task);
2464        let client_slot = Arc::new(Mutex::new(None));
2465        let client_slot_callback = Arc::clone(&client_slot);
2466        let states = Arc::new(Mutex::new(Vec::new()));
2467        let states_callback = Arc::clone(&states);
2468        client.state_sink = Some(SocketStateSink::new(move |state| {
2469            states_callback.lock().push(state);
2470            drop(client_slot_callback.lock().take());
2471        }));
2472        let handle = client.reconnect_handle();
2473        let controller_notify = Arc::clone(&handle.controller_notify);
2474        *client_slot.lock() = Some(client);
2475        let (result_tx, result_rx) = std::sync::mpsc::channel();
2476        std::thread::spawn(move || result_tx.send(handle.request_reconnect()).unwrap());
2477
2478        assert_eq!(
2479            recv_rendezvous(result_rx, "reconnect request").await,
2480            ReconnectRequestOutcome::Accepted
2481        );
2482        tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
2483            .await
2484            .expect("accepted request should notify before deferred controller abort");
2485        wait_until_async(|| async { controller_abort.is_finished() }, TEST_TIMEOUT).await;
2486
2487        assert!(client_slot.lock().is_none());
2488        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2489    }
2490
2491    #[rstest]
2492    #[tokio::test]
2493    async fn test_manual_reconnect_uses_controller_path() {
2494        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2495        let port = listener.local_addr().unwrap().port();
2496        let (first_accepted_tx, first_accepted_rx) = oneshot::channel();
2497        let (second_accepted_tx, second_accepted_rx) = oneshot::channel();
2498        let (payload_tx, payload_rx) = oneshot::channel();
2499        let (release_tx, release_rx) = oneshot::channel();
2500
2501        let server = task::spawn(async move {
2502            let (_first, _) = listener.accept().await.unwrap();
2503            first_accepted_tx.send(()).unwrap();
2504
2505            let (mut second, _) = listener.accept().await.unwrap();
2506            second_accepted_tx.send(()).unwrap();
2507            let mut payload = [0_u8; 8];
2508            second.read_exact(&mut payload).await.unwrap();
2509            payload_tx.send(payload).unwrap();
2510            let _ = release_rx.await;
2511        });
2512        let states = Arc::new(Mutex::new(Vec::new()));
2513        let states_callback = Arc::clone(&states);
2514        let sink = SocketStateSink::new(move |state| {
2515            states_callback.lock().push(state);
2516        });
2517        let callback_count = Arc::new(AtomicUsize::new(0));
2518        let callback_count_clone = Arc::clone(&callback_count);
2519        let post_reconnection = Arc::new(move || {
2520            callback_count_clone.fetch_add(1, Ordering::SeqCst);
2521        });
2522        let client = SocketClient::builder()
2523            .config(reconnect_test_config(port))
2524            .post_reconnection(post_reconnection)
2525            .state_sink(sink)
2526            .connect()
2527            .await
2528            .unwrap();
2529        first_accepted_rx.await.unwrap();
2530        let handle = client.reconnect_handle();
2531
2532        assert!(client.request_reconnect());
2533        assert_eq!(
2534            handle.request_reconnect(),
2535            ReconnectRequestOutcome::AlreadyReconnecting
2536        );
2537        assert!(!client.request_reconnect());
2538        assert_eq!(
2539            *states.lock(),
2540            vec![SocketState::Connected, SocketState::Disconnected]
2541        );
2542
2543        tokio::time::timeout(TEST_TIMEOUT, second_accepted_rx)
2544            .await
2545            .expect("controller should establish a replacement connection")
2546            .unwrap();
2547        wait_until_async(
2548            || async { client.is_active() && callback_count.load(Ordering::SeqCst) == 1 },
2549            TEST_TIMEOUT,
2550        )
2551        .await;
2552        client.send_bytes(b"manual".to_vec()).await.unwrap();
2553
2554        assert_eq!(
2555            tokio::time::timeout(TEST_TIMEOUT, payload_rx)
2556                .await
2557                .expect("replacement connection should receive the framed payload")
2558                .unwrap(),
2559            *b"manual\r\n"
2560        );
2561        assert_eq!(callback_count.load(Ordering::SeqCst), 1);
2562        assert_eq!(
2563            *states.lock(),
2564            vec![
2565                SocketState::Connected,
2566                SocketState::Disconnected,
2567                SocketState::Connected,
2568            ]
2569        );
2570
2571        client.close().await;
2572        release_tx.send(()).unwrap();
2573        server.await.unwrap();
2574        assert_eq!(callback_count.load(Ordering::SeqCst), 1);
2575        assert_eq!(states.lock().len(), 3);
2576    }
2577
2578    #[rstest]
2579    #[case(ConnectionMode::Disconnect)]
2580    #[case(ConnectionMode::Closed)]
2581    #[tokio::test]
2582    async fn test_send_bytes_rejects_terminal_state(#[case] mode: ConnectionMode) {
2583        let connection_state = Arc::new(AtomicU8::new(mode.as_u8()));
2584        let state_notify = Arc::new(tokio::sync::Notify::new());
2585        let controller_task = tokio::spawn(std::future::pending::<()>());
2586        let client = test_socket_client(connection_state, state_notify, controller_task);
2587
2588        let result = client.send_bytes(b"terminal".to_vec()).await;
2589
2590        assert!(matches!(result, Err(SendError::Closed)));
2591    }
2592
2593    #[rstest]
2594    #[tokio::test(start_paused = true)]
2595    async fn test_close_sets_closed_after_controller_was_aborted() {
2596        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2597        let state_notify = Arc::new(tokio::sync::Notify::new());
2598
2599        let controller_task = tokio::spawn(std::future::pending::<()>());
2600        controller_task.abort();
2601        let client = test_socket_client(connection_state, state_notify, controller_task);
2602
2603        client.close_with_timeout(Duration::from_millis(1)).await;
2604
2605        assert!(client.is_closed());
2606    }
2607
2608    #[rstest]
2609    #[tokio::test]
2610    async fn test_reconnect_exhaustion_closes_client() {
2611        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2612        let port = listener.local_addr().unwrap().port();
2613        let (accepted_tx, accepted_rx) = oneshot::channel();
2614        let (release_tx, release_rx) = oneshot::channel();
2615
2616        let server = task::spawn(async move {
2617            let (socket, _) = listener.accept().await.unwrap();
2618            let _ = accepted_tx.send(());
2619            let _ = release_rx.await;
2620            drop(socket);
2621            drop(listener);
2622        });
2623
2624        let config = SocketConfig {
2625            url: format!("127.0.0.1:{port}"),
2626            mode: Mode::Plain,
2627            suffix: b"\r\n".to_vec(),
2628            message_handler: None,
2629            heartbeat: None,
2630            connect_timeout_ms: Some(100),
2631            reconnect_delay_initial_ms: Some(1),
2632            reconnect_delay_max_ms: Some(1),
2633            reconnect_backoff_factor: Some(1.0),
2634            reconnect_jitter_ms: Some(0),
2635            connection_max_retries: Some(1),
2636            reconnect_max_attempts: Some(1),
2637            heartbeat_timeout_secs: None,
2638            certs_dir: None,
2639        };
2640
2641        let states = Arc::new(Mutex::new(Vec::new()));
2642        let states_callback = Arc::clone(&states);
2643        let sink = SocketStateSink::new(move |state| {
2644            states_callback.lock().push(state);
2645        });
2646        let client = SocketClient::builder()
2647            .config(config)
2648            .state_sink(sink)
2649            .connect()
2650            .await
2651            .unwrap();
2652        accepted_rx.await.unwrap();
2653        release_tx.send(()).unwrap();
2654
2655        wait_until_async(|| async { client.is_closed() }, Duration::from_secs(5)).await;
2656        assert!(client.is_closed());
2657        assert_eq!(
2658            *states.lock(),
2659            vec![SocketState::Connected, SocketState::Disconnected]
2660        );
2661
2662        client.close().await;
2663        assert!(client.is_closed());
2664        assert_eq!(states.lock().len(), 2);
2665        server.await.unwrap();
2666    }
2667
2668    #[rstest]
2669    #[tokio::test]
2670    async fn test_drop_suppresses_socket_state_event() {
2671        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2672        let port = listener.local_addr().unwrap().port();
2673        let (accepted_tx, accepted_rx) = oneshot::channel();
2674
2675        let server = task::spawn(async move {
2676            let (_socket, _) = listener.accept().await.unwrap();
2677            let _ = accepted_tx.send(());
2678            std::future::pending::<()>().await;
2679        });
2680
2681        let config = SocketConfig {
2682            url: format!("127.0.0.1:{port}"),
2683            mode: Mode::Plain,
2684            suffix: b"\r\n".to_vec(),
2685            message_handler: None,
2686            heartbeat: None,
2687            connect_timeout_ms: Some(100),
2688            reconnect_delay_initial_ms: Some(1),
2689            reconnect_delay_max_ms: Some(1),
2690            reconnect_backoff_factor: Some(1.0),
2691            reconnect_jitter_ms: Some(0),
2692            connection_max_retries: Some(1),
2693            reconnect_max_attempts: Some(1),
2694            heartbeat_timeout_secs: None,
2695            certs_dir: None,
2696        };
2697        let states = Arc::new(Mutex::new(Vec::new()));
2698        let states_callback = Arc::clone(&states);
2699        let sink = SocketStateSink::new(move |state| {
2700            states_callback.lock().push(state);
2701        });
2702        let client = SocketClient::builder()
2703            .config(config)
2704            .state_sink(sink)
2705            .connect()
2706            .await
2707            .unwrap();
2708        accepted_rx.await.unwrap();
2709
2710        drop(client);
2711        sleep(Duration::from_millis(25)).await;
2712
2713        assert_eq!(*states.lock(), vec![SocketState::Connected]);
2714        server.abort();
2715    }
2716
2717    #[rstest]
2718    #[tokio::test]
2719    async fn test_graceful_close_is_idempotent() {
2720        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2721        let port = listener.local_addr().unwrap().port();
2722        let (accepted_tx, accepted_rx) = oneshot::channel();
2723
2724        let server = task::spawn(async move {
2725            let (_socket, _) = listener.accept().await.unwrap();
2726            let _ = accepted_tx.send(());
2727            std::future::pending::<()>().await;
2728        });
2729        let config = SocketConfig {
2730            url: format!("127.0.0.1:{port}"),
2731            mode: Mode::Plain,
2732            suffix: b"\r\n".to_vec(),
2733            message_handler: None,
2734            heartbeat: None,
2735            connect_timeout_ms: None,
2736            reconnect_delay_initial_ms: None,
2737            reconnect_delay_max_ms: None,
2738            reconnect_backoff_factor: None,
2739            reconnect_jitter_ms: None,
2740            connection_max_retries: None,
2741            reconnect_max_attempts: None,
2742            heartbeat_timeout_secs: None,
2743            certs_dir: None,
2744        };
2745        let client = SocketClient::builder()
2746            .config(config)
2747            .connect()
2748            .await
2749            .unwrap();
2750        accepted_rx.await.unwrap();
2751
2752        client.close().await;
2753        client.close().await;
2754
2755        assert!(client.is_closed());
2756        server.abort();
2757    }
2758
2759    #[rstest]
2760    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
2761    async fn test_read_loop_drops_remaining_frames_after_session_replaced() {
2762        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2763        let read_fence = ReadSessionFence::new();
2764        let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
2765        let received_clone = Arc::clone(&received);
2766        let handler_release = Arc::new((Mutex::new(false), Condvar::new()));
2767        let handler_release_guard = CondvarReleaseGuard::new(Arc::clone(&handler_release));
2768        let handler_release_clone = Arc::clone(&handler_release);
2769        let (handler_entered_tx, handler_entered_rx) = std::sync::mpsc::channel();
2770        let handler: TcpMessageHandler = Arc::new(move |data| {
2771            received_clone.lock().push(data.to_vec());
2772            handler_entered_tx.send(()).unwrap();
2773            let (lock, condvar) = handler_release_clone.as_ref();
2774            let mut released = lock.lock();
2775
2776            while !*released {
2777                condvar.wait(&mut released);
2778            }
2779        });
2780        let reader = ScriptedReader {
2781            first: Some(b"first\r\nsecond\r\n".to_vec()),
2782            remainder: Arc::new(Mutex::new(None)),
2783            pending_tx: None,
2784            waker: Arc::new(Mutex::new(None)),
2785        };
2786        let read_task = SocketClientInner::spawn_read_task(
2787            Arc::clone(&connection_state),
2788            read_fence.clone(),
2789            reader,
2790            Some(handler),
2791            b"\r\n".to_vec(),
2792            None,
2793        );
2794
2795        recv_rendezvous(handler_entered_rx, "socket first-handler entry").await;
2796        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2797        read_fence.invalidate();
2798        read_task.abort();
2799        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
2800
2801        handler_release_guard.release();
2802        await_task_termination(read_task, "old socket read task", true).await;
2803
2804        assert_eq!(received.lock().as_slice(), &[b"first".to_vec()]);
2805    }
2806
2807    #[rstest]
2808    #[tokio::test(start_paused = true)]
2809    async fn test_read_loop_drops_partial_old_session_frame() {
2810        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2811        let old_read_fence = ReadSessionFence::new();
2812        let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
2813        let received_clone = Arc::clone(&received);
2814        let handler: TcpMessageHandler =
2815            Arc::new(move |data| received_clone.lock().push(data.to_vec()));
2816        let (pending_tx, pending_rx) = oneshot::channel();
2817        let remainder = Arc::new(Mutex::new(None));
2818        let waker = Arc::new(Mutex::new(None));
2819        let reader = ScriptedReader {
2820            first: Some(b"old".to_vec()),
2821            remainder: Arc::clone(&remainder),
2822            pending_tx: Some(pending_tx),
2823            waker: Arc::clone(&waker),
2824        };
2825
2826        let read_task = SocketClientInner::spawn_read_task(
2827            Arc::clone(&connection_state),
2828            old_read_fence.clone(),
2829            reader,
2830            Some(handler),
2831            b"\r\n".to_vec(),
2832            None,
2833        );
2834
2835        pending_rx.await.unwrap();
2836        connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2837        old_read_fence.invalidate();
2838        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
2839        *remainder.lock() = Some(b"\r\nnew\r\n".to_vec());
2840        waker.lock().take().unwrap().wake();
2841        read_task.await.unwrap();
2842
2843        let fresh_reader = ScriptedReader {
2844            first: Some(b"new\r\n".to_vec()),
2845            remainder: Arc::new(Mutex::new(None)),
2846            pending_tx: None,
2847            waker: Arc::new(Mutex::new(None)),
2848        };
2849        let fresh_connection_state = Arc::clone(&connection_state);
2850        let fresh_received = Arc::clone(&received);
2851        let fresh_handler: TcpMessageHandler = Arc::new(move |data| {
2852            fresh_received.lock().push(data.to_vec());
2853            fresh_connection_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2854        });
2855        SocketClientInner::spawn_read_task(
2856            Arc::clone(&connection_state),
2857            ReadSessionFence::new(),
2858            fresh_reader,
2859            Some(fresh_handler),
2860            b"\r\n".to_vec(),
2861            None,
2862        )
2863        .await
2864        .unwrap();
2865
2866        assert_eq!(received.lock().as_slice(), &[b"new".to_vec()]);
2867    }
2868
2869    #[rstest]
2870    #[tokio::test(start_paused = true)]
2871    async fn test_read_loop_stops_when_first_handler_ends_session() {
2872        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2873        let handler_state = Arc::clone(&connection_state);
2874        let received = Arc::new(Mutex::new(Vec::<Vec<u8>>::new()));
2875        let received_clone = Arc::clone(&received);
2876        let handler: TcpMessageHandler = Arc::new(move |data| {
2877            received_clone.lock().push(data.to_vec());
2878            handler_state.store(ConnectionMode::Reconnect.as_u8(), Ordering::SeqCst);
2879        });
2880        let (pending_tx, _pending_rx) = oneshot::channel();
2881        let reader = ScriptedReader {
2882            first: Some(b"first\r\nsecond\r\n".to_vec()),
2883            remainder: Arc::new(Mutex::new(None)),
2884            pending_tx: Some(pending_tx),
2885            waker: Arc::new(Mutex::new(None)),
2886        };
2887
2888        SocketClientInner::spawn_read_task(
2889            connection_state,
2890            ReadSessionFence::new(),
2891            reader,
2892            Some(handler),
2893            b"\r\n".to_vec(),
2894            None,
2895        )
2896        .await
2897        .unwrap();
2898
2899        assert_eq!(received.lock().as_slice(), &[b"first".to_vec()]);
2900    }
2901
2902    #[tokio::test]
2903    async fn test_read_loop_closes_when_unframed_buffer_exceeds_limit() {
2904        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2905        let read_after_limit = Arc::new(AtomicBool::new(false));
2906        let reader = LimitProbeReader {
2907            remaining: MAX_READ_BUFFER_BYTES + 1,
2908            read_after_limit: Arc::clone(&read_after_limit),
2909        };
2910
2911        SocketClientInner::run_read_loop(
2912            connection_state,
2913            ReadSessionFence::new(),
2914            reader,
2915            None,
2916            b"\r\n".to_vec(),
2917            None,
2918            Duration::from_secs(1),
2919        )
2920        .await;
2921
2922        assert!(!read_after_limit.load(Ordering::SeqCst));
2923    }
2924
2925    #[rstest]
2926    #[tokio::test(start_paused = true)]
2927    async fn test_stalled_socket_write_sends_reconnect_replay_before_buffer() {
2928        type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
2929
2930        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
2931        let state_notify = Arc::new(tokio::sync::Notify::new());
2932        let (stream, _non_reading_peer) = tokio::io::duplex(1);
2933        let (pending_tx, pending_rx) = oneshot::channel();
2934        let writer: TestWriter = Box::pin(BackpressuredWriter {
2935            stream,
2936            pending_tx: Some(pending_tx),
2937        });
2938        let (writer_tx, writer_rx) =
2939            tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
2940        let states = Arc::new(Mutex::new(Vec::new()));
2941        let states_callback = Arc::clone(&states);
2942        let sink = SocketStateSink::new(move |state| {
2943            states_callback.lock().push(state);
2944        });
2945        let write_task = SocketClientInner::spawn_write_task(
2946            Arc::clone(&connection_state),
2947            Arc::clone(&state_notify),
2948            writer,
2949            writer_rx,
2950            b"\r\n".to_vec(),
2951            Some(sink),
2952        );
2953
2954        writer_tx
2955            .send(WriterCommand::Send(Bytes::from_static(b"complete-message")))
2956            .unwrap();
2957        pending_rx.await.unwrap();
2958
2959        let recorded = Arc::new(Mutex::new(Vec::new()));
2960        let new_writer: TestWriter = Box::pin(RecordingWriter {
2961            bytes: Arc::clone(&recorded),
2962        });
2963        let (update_tx, update_rx) = oneshot::channel();
2964        writer_tx
2965            .send(WriterCommand::UpdateWithReplay(
2966                new_writer,
2967                vec![Bytes::from_static(b"authentication")],
2968                update_tx,
2969            ))
2970            .unwrap();
2971
2972        tokio::time::advance(Duration::from_secs(GRACEFUL_SHUTDOWN_TIMEOUT_SECS)).await;
2973        yield_now().await;
2974
2975        assert!(
2976            tokio::time::timeout(Duration::from_secs(10), update_rx)
2977                .await
2978                .expect("writer update should not remain queued behind a stalled write")
2979                .unwrap(),
2980            "writer should report successful replay"
2981        );
2982        assert_eq!(
2983            ConnectionMode::from_atomic(&connection_state),
2984            ConnectionMode::Reconnect
2985        );
2986        assert_eq!(
2987            recorded.lock().as_slice(),
2988            b"authentication\r\ncomplete-message\r\n"
2989        );
2990
2991        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
2992        state_notify.notify_waiters();
2993        drop(writer_tx);
2994        write_task.await.unwrap();
2995
2996        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
2997    }
2998
2999    #[rstest]
3000    #[tokio::test(start_paused = true)]
3001    async fn test_send_after_writer_update_drains_when_reconnect_completes() {
3002        type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3003
3004        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3005        let state_notify = Arc::new(tokio::sync::Notify::new());
3006        let initial_writer: TestWriter = Box::pin(RecordingWriter {
3007            bytes: Arc::new(Mutex::new(Vec::new())),
3008        });
3009        let (writer_tx, writer_rx) =
3010            tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3011        let write_task = SocketClientInner::spawn_write_task(
3012            Arc::clone(&connection_state),
3013            Arc::clone(&state_notify),
3014            initial_writer,
3015            writer_rx,
3016            b"\r\n".to_vec(),
3017            None,
3018        );
3019
3020        let recorded = Arc::new(Mutex::new(Vec::new()));
3021        let new_writer: TestWriter = Box::pin(RecordingWriter {
3022            bytes: Arc::clone(&recorded),
3023        });
3024        let (update_tx, update_rx) = oneshot::channel();
3025        writer_tx
3026            .send(WriterCommand::Update(new_writer, update_tx))
3027            .unwrap();
3028        assert!(update_rx.await.unwrap());
3029
3030        writer_tx
3031            .send(WriterCommand::Send(Bytes::from_static(b"late")))
3032            .unwrap();
3033        yield_now().await;
3034
3035        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3036        state_notify.notify_waiters();
3037        writer_tx
3038            .send(WriterCommand::Send(Bytes::from_static(b"new")))
3039            .unwrap();
3040        tokio::time::advance(Duration::from_millis(
3041            CONNECTION_STATE_CHECK_INTERVAL_MS * 2,
3042        ))
3043        .await;
3044        yield_now().await;
3045
3046        let actual = recorded.lock().clone();
3047        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3048        state_notify.notify_waiters();
3049        drop(writer_tx);
3050        write_task.await.unwrap();
3051
3052        assert_eq!(actual, b"late\r\nnew\r\n");
3053    }
3054
3055    #[rstest]
3056    #[case::latest_wins(
3057        &["subscription-a", "subscription-b"],
3058        &["subscription-b"],
3059        b"subscription-b\r\n",
3060    )]
3061    #[case::different_replay_drains(
3062        &["subscription-b"],
3063        &["subscription-a"],
3064        b"subscription-a\r\nsubscription-b\r\n",
3065    )]
3066    #[tokio::test]
3067    async fn test_send_or_replay_buffering(
3068        #[case] buffered: &[&str],
3069        #[case] replay: &[&str],
3070        #[case] expected: &[u8],
3071    ) {
3072        type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3073
3074        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3075        let state_notify = Arc::new(tokio::sync::Notify::new());
3076        let initial_writer: TestWriter = Box::pin(RecordingWriter {
3077            bytes: Arc::new(Mutex::new(Vec::new())),
3078        });
3079        let (writer_tx, writer_rx) =
3080            tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3081        let write_task = SocketClientInner::spawn_write_task(
3082            Arc::clone(&connection_state),
3083            Arc::clone(&state_notify),
3084            initial_writer,
3085            writer_rx,
3086            b"\r\n".to_vec(),
3087            None,
3088        );
3089
3090        for data in buffered {
3091            writer_tx
3092                .send(WriterCommand::SendOrReplay {
3093                    key: 7,
3094                    data: Bytes::copy_from_slice(data.as_bytes()),
3095                })
3096                .unwrap();
3097        }
3098        let recorded = Arc::new(Mutex::new(Vec::new()));
3099        let new_writer: TestWriter = Box::pin(RecordingWriter {
3100            bytes: Arc::clone(&recorded),
3101        });
3102        let (update_tx, update_rx) = oneshot::channel();
3103        writer_tx
3104            .send(WriterCommand::UpdateWithReplay(
3105                new_writer,
3106                replay
3107                    .iter()
3108                    .map(|data| Bytes::copy_from_slice(data.as_bytes()))
3109                    .collect(),
3110                update_tx,
3111            ))
3112            .unwrap();
3113
3114        assert!(update_rx.await.unwrap());
3115        assert_eq!(recorded.lock().as_slice(), expected);
3116        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3117        state_notify.notify_waiters();
3118        drop(writer_tx);
3119        write_task.await.unwrap();
3120    }
3121
3122    #[rstest]
3123    #[tokio::test(start_paused = true)]
3124    async fn test_active_reconnect_buffer_write_failure_reconnects_and_retries() {
3125        type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3126
3127        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Reconnect.as_u8()));
3128        let state_notify = Arc::new(tokio::sync::Notify::new());
3129        let initial_writer: TestWriter = Box::pin(RecordingWriter {
3130            bytes: Arc::new(Mutex::new(Vec::new())),
3131        });
3132        let (writer_tx, writer_rx) =
3133            tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3134        let states = Arc::new(Mutex::new(Vec::new()));
3135        let states_callback = Arc::clone(&states);
3136        let sink = SocketStateSink::new(move |state| {
3137            states_callback.lock().push(state);
3138        });
3139        let write_task = SocketClientInner::spawn_write_task(
3140            Arc::clone(&connection_state),
3141            Arc::clone(&state_notify),
3142            initial_writer,
3143            writer_rx,
3144            b"\r\n".to_vec(),
3145            Some(sink),
3146        );
3147
3148        let (update_tx, update_rx) = oneshot::channel();
3149        let failing_writer: TestWriter = Box::pin(FailingWriter);
3150        writer_tx
3151            .send(WriterCommand::Update(failing_writer, update_tx))
3152            .unwrap();
3153        assert!(update_rx.await.unwrap());
3154
3155        writer_tx
3156            .send(WriterCommand::Send(Bytes::from_static(b"late")))
3157            .unwrap();
3158        yield_now().await;
3159        connection_state.store(ConnectionMode::Active.as_u8(), Ordering::SeqCst);
3160        state_notify.notify_waiters();
3161        tokio::time::advance(Duration::from_millis(
3162            CONNECTION_STATE_CHECK_INTERVAL_MS * 2,
3163        ))
3164        .await;
3165        yield_now().await;
3166
3167        assert_eq!(
3168            ConnectionMode::from_atomic(&connection_state),
3169            ConnectionMode::Reconnect
3170        );
3171        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
3172
3173        let recorded = Arc::new(Mutex::new(Vec::new()));
3174        let replacement: TestWriter = Box::pin(RecordingWriter {
3175            bytes: Arc::clone(&recorded),
3176        });
3177        let (retry_tx, retry_rx) = oneshot::channel();
3178        writer_tx
3179            .send(WriterCommand::Update(replacement, retry_tx))
3180            .unwrap();
3181        assert!(retry_rx.await.unwrap());
3182
3183        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3184        state_notify.notify_waiters();
3185        drop(writer_tx);
3186        write_task.await.unwrap();
3187
3188        assert_eq!(recorded.lock().as_slice(), b"late\r\n");
3189    }
3190
3191    #[rstest]
3192    #[tokio::test(start_paused = true)]
3193    async fn test_failed_send_or_replay_is_not_duplicated_after_replay() {
3194        type TestWriter = Pin<Box<dyn AsyncWrite + Send>>;
3195
3196        let connection_state = Arc::new(AtomicU8::new(ConnectionMode::Active.as_u8()));
3197        let state_notify = Arc::new(tokio::sync::Notify::new());
3198        let (stream, _non_reading_peer) = tokio::io::duplex(1);
3199        let (pending_tx, pending_rx) = oneshot::channel();
3200        let writer: TestWriter = Box::pin(BackpressuredWriter {
3201            stream,
3202            pending_tx: Some(pending_tx),
3203        });
3204        let (writer_tx, writer_rx) =
3205            tokio::sync::mpsc::unbounded_channel::<WriterCommand<TestWriter>>();
3206        let write_task = SocketClientInner::spawn_write_task(
3207            Arc::clone(&connection_state),
3208            Arc::clone(&state_notify),
3209            writer,
3210            writer_rx,
3211            b"\r\n".to_vec(),
3212            None,
3213        );
3214        let subscription = Bytes::from_static(b"subscription");
3215
3216        writer_tx
3217            .send(WriterCommand::SendOrReplay {
3218                key: 7,
3219                data: subscription.clone(),
3220            })
3221            .unwrap();
3222        pending_rx.await.unwrap();
3223
3224        let recorded = Arc::new(Mutex::new(Vec::new()));
3225        let new_writer: TestWriter = Box::pin(RecordingWriter {
3226            bytes: Arc::clone(&recorded),
3227        });
3228        let (update_tx, update_rx) = oneshot::channel();
3229        writer_tx
3230            .send(WriterCommand::UpdateWithReplay(
3231                new_writer,
3232                vec![subscription],
3233                update_tx,
3234            ))
3235            .unwrap();
3236
3237        tokio::time::advance(Duration::from_secs(WRITE_TIMEOUT_SECS)).await;
3238        yield_now().await;
3239        assert!(update_rx.await.unwrap());
3240        assert_eq!(recorded.lock().as_slice(), b"subscription\r\n");
3241
3242        connection_state.store(ConnectionMode::Closed.as_u8(), Ordering::SeqCst);
3243        state_notify.notify_waiters();
3244        drop(writer_tx);
3245        write_task.await.unwrap();
3246    }
3247
3248    #[rstest]
3249    #[tokio::test]
3250    async fn test_connect_url_rejects_invalid_reconnect_backoff_before_connect() {
3251        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
3252        listener.set_nonblocking(true).unwrap();
3253        let port = listener.local_addr().unwrap().port();
3254        let config = SocketConfig {
3255            url: format!("127.0.0.1:{port}"),
3256            mode: Mode::Plain,
3257            suffix: b"\r\n".to_vec(),
3258            message_handler: None,
3259            heartbeat: None,
3260            connect_timeout_ms: Some(1_000),
3261            reconnect_delay_initial_ms: Some(50),
3262            reconnect_delay_max_ms: Some(100),
3263            reconnect_backoff_factor: Some(100.1),
3264            reconnect_jitter_ms: Some(0),
3265            connection_max_retries: Some(1),
3266            reconnect_max_attempts: None,
3267            heartbeat_timeout_secs: None,
3268            certs_dir: None,
3269        };
3270
3271        let error = match SocketClientInner::connect_url(config, None).await {
3272            Ok(_) => panic!("invalid reconnect backoff should be rejected"),
3273            Err(e) => e,
3274        };
3275
3276        assert!(
3277            error.to_string().contains("factor"),
3278            "error should mention the invalid factor, was: {error}"
3279        );
3280        assert_eq!(
3281            listener.accept().unwrap_err().kind(),
3282            std::io::ErrorKind::WouldBlock,
3283            "invalid reconnect backoff must be rejected before connecting"
3284        );
3285    }
3286
3287    #[rstest]
3288    #[tokio::test]
3289    async fn test_reconnect_then_close() {
3290        // Bind an ephemeral port
3291        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3292        let port = listener.local_addr().unwrap().port();
3293
3294        // Server task: accept one connection and then drop it
3295        let server = task::spawn(async move {
3296            if let Ok((mut sock, _)) = listener.accept().await {
3297                drop(sock.shutdown());
3298            }
3299            // Keep listener alive briefly to avoid premature exit
3300            sleep(Duration::from_secs(1)).await;
3301        });
3302
3303        // Configure client with a short reconnect backoff
3304        let config = SocketConfig {
3305            url: format!("127.0.0.1:{port}"),
3306            mode: Mode::Plain,
3307            suffix: b"\r\n".to_vec(),
3308            message_handler: None,
3309            heartbeat: None,
3310            connect_timeout_ms: Some(1_000),
3311            reconnect_delay_initial_ms: Some(50),
3312            reconnect_delay_max_ms: Some(100),
3313            reconnect_backoff_factor: Some(1.0),
3314            reconnect_jitter_ms: Some(0),
3315            connection_max_retries: Some(1),
3316            reconnect_max_attempts: None,
3317            heartbeat_timeout_secs: None,
3318            certs_dir: None,
3319        };
3320
3321        // Connect client (handler=None)
3322        let client = SocketClient::builder()
3323            .config(config.clone())
3324            .connect()
3325            .await
3326            .unwrap();
3327
3328        // Wait for client to detect dropped connection and enter reconnect state
3329        wait_until_async(
3330            || async { client.is_reconnecting() },
3331            Duration::from_secs(2),
3332        )
3333        .await;
3334
3335        // Now close the client
3336        client.close().await;
3337        assert!(client.is_closed());
3338        server.abort();
3339    }
3340
3341    #[rstest]
3342    #[tokio::test]
3343    async fn test_reconnect_state_flips_when_reader_stops() {
3344        // Bind an ephemeral port and accept a single connection which we immediately close.
3345        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3346        let port = listener.local_addr().unwrap().port();
3347
3348        let server = task::spawn(async move {
3349            if let Ok((sock, _)) = listener.accept().await {
3350                drop(sock);
3351            }
3352            // Give the client a moment to observe the closed connection.
3353            sleep(Duration::from_millis(50)).await;
3354        });
3355
3356        let config = SocketConfig {
3357            url: format!("127.0.0.1:{port}"),
3358            mode: Mode::Plain,
3359            suffix: b"\r\n".to_vec(),
3360            message_handler: None,
3361            heartbeat: None,
3362            connect_timeout_ms: Some(1_000),
3363            reconnect_delay_initial_ms: Some(50),
3364            reconnect_delay_max_ms: Some(100),
3365            reconnect_backoff_factor: Some(1.0),
3366            reconnect_jitter_ms: Some(0),
3367            connection_max_retries: Some(1),
3368            reconnect_max_attempts: None,
3369            heartbeat_timeout_secs: None,
3370            certs_dir: None,
3371        };
3372
3373        let client = SocketClient::builder()
3374            .config(config)
3375            .connect()
3376            .await
3377            .unwrap();
3378
3379        wait_until_async(
3380            || async { client.is_reconnecting() },
3381            Duration::from_secs(2),
3382        )
3383        .await;
3384
3385        client.close().await;
3386        server.abort();
3387    }
3388
3389    #[rstest]
3390    fn test_parse_socket_url_raw_address() {
3391        // Raw socket address with TLS mode
3392        let (socket_addr, request_url) =
3393            SocketClientInner::parse_socket_url("example.com:6130", Mode::Tls).unwrap();
3394        assert_eq!(socket_addr, "example.com:6130");
3395        assert_eq!(request_url, "wss://example.com:6130");
3396
3397        // Raw socket address with Plain mode
3398        let (socket_addr, request_url) =
3399            SocketClientInner::parse_socket_url("localhost:8080", Mode::Plain).unwrap();
3400        assert_eq!(socket_addr, "localhost:8080");
3401        assert_eq!(request_url, "ws://localhost:8080");
3402    }
3403
3404    #[rstest]
3405    fn test_parse_socket_url_with_scheme() {
3406        // Full URL with wss scheme
3407        let (socket_addr, request_url) =
3408            SocketClientInner::parse_socket_url("wss://example.com:443/path", Mode::Tls).unwrap();
3409        assert_eq!(socket_addr, "example.com:443");
3410        assert_eq!(request_url, "wss://example.com:443/path");
3411
3412        // Full URL with ws scheme
3413        let (socket_addr, request_url) =
3414            SocketClientInner::parse_socket_url("ws://localhost:8080", Mode::Plain).unwrap();
3415        assert_eq!(socket_addr, "localhost:8080");
3416        assert_eq!(request_url, "ws://localhost:8080");
3417    }
3418
3419    #[rstest]
3420    fn test_parse_socket_url_default_ports() {
3421        // wss without explicit port defaults to 443
3422        let (socket_addr, _) =
3423            SocketClientInner::parse_socket_url("wss://example.com", Mode::Tls).unwrap();
3424        assert_eq!(socket_addr, "example.com:443");
3425
3426        // ws without explicit port defaults to 80
3427        let (socket_addr, _) =
3428            SocketClientInner::parse_socket_url("ws://example.com", Mode::Plain).unwrap();
3429        assert_eq!(socket_addr, "example.com:80");
3430
3431        // https defaults to 443
3432        let (socket_addr, _) =
3433            SocketClientInner::parse_socket_url("https://example.com", Mode::Tls).unwrap();
3434        assert_eq!(socket_addr, "example.com:443");
3435
3436        // http defaults to 80
3437        let (socket_addr, _) =
3438            SocketClientInner::parse_socket_url("http://example.com", Mode::Plain).unwrap();
3439        assert_eq!(socket_addr, "example.com:80");
3440    }
3441
3442    #[rstest]
3443    fn test_parse_socket_url_unknown_scheme_uses_mode() {
3444        // Unknown scheme defaults to mode-based port
3445        let (socket_addr, _) =
3446            SocketClientInner::parse_socket_url("custom://example.com", Mode::Tls).unwrap();
3447        assert_eq!(socket_addr, "example.com:443");
3448
3449        let (socket_addr, _) =
3450            SocketClientInner::parse_socket_url("custom://example.com", Mode::Plain).unwrap();
3451        assert_eq!(socket_addr, "example.com:80");
3452    }
3453
3454    #[rstest]
3455    fn test_parse_socket_url_ipv6() {
3456        // IPv6 address with port
3457        let (socket_addr, request_url) =
3458            SocketClientInner::parse_socket_url("[::1]:8080", Mode::Plain).unwrap();
3459        assert_eq!(socket_addr, "[::1]:8080");
3460        assert_eq!(request_url, "ws://[::1]:8080");
3461
3462        // IPv6 in URL
3463        let (socket_addr, _) =
3464            SocketClientInner::parse_socket_url("ws://[::1]:8080", Mode::Plain).unwrap();
3465        assert_eq!(socket_addr, "[::1]:8080");
3466    }
3467
3468    #[rstest]
3469    #[tokio::test]
3470    async fn test_url_parsing_raw_socket_address() {
3471        // Test that raw socket addresses (host:port) work correctly
3472        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3473        let port = listener.local_addr().unwrap().port();
3474
3475        let server = task::spawn(async move {
3476            if let Ok((sock, _)) = listener.accept().await {
3477                drop(sock);
3478            }
3479            sleep(Duration::from_millis(50)).await;
3480        });
3481
3482        let config = SocketConfig {
3483            url: format!("127.0.0.1:{port}"), // Raw socket address format
3484            mode: Mode::Plain,
3485            suffix: b"\r\n".to_vec(),
3486            message_handler: None,
3487            heartbeat: None,
3488            connect_timeout_ms: Some(1_000),
3489            reconnect_delay_initial_ms: Some(50),
3490            reconnect_delay_max_ms: Some(100),
3491            reconnect_backoff_factor: Some(1.0),
3492            reconnect_jitter_ms: Some(0),
3493            connection_max_retries: Some(1),
3494            reconnect_max_attempts: None,
3495            heartbeat_timeout_secs: None,
3496            certs_dir: None,
3497        };
3498
3499        // Should successfully connect with raw socket address
3500        let client = SocketClient::builder().config(config).connect().await;
3501        assert!(
3502            client.is_ok(),
3503            "Client should connect with raw socket address format"
3504        );
3505
3506        if let Ok(client) = client {
3507            client.close().await;
3508        }
3509        server.abort();
3510    }
3511
3512    #[rstest]
3513    #[tokio::test]
3514    async fn test_url_parsing_with_scheme() {
3515        // Test that URLs with schemes also work
3516        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3517        let port = listener.local_addr().unwrap().port();
3518
3519        let server = task::spawn(async move {
3520            if let Ok((sock, _)) = listener.accept().await {
3521                drop(sock);
3522            }
3523            sleep(Duration::from_millis(50)).await;
3524        });
3525
3526        let config = SocketConfig {
3527            url: format!("ws://127.0.0.1:{port}"), // URL with scheme
3528            mode: Mode::Plain,
3529            suffix: b"\r\n".to_vec(),
3530            message_handler: None,
3531            heartbeat: None,
3532            connect_timeout_ms: Some(1_000),
3533            reconnect_delay_initial_ms: Some(50),
3534            reconnect_delay_max_ms: Some(100),
3535            reconnect_backoff_factor: Some(1.0),
3536            reconnect_jitter_ms: Some(0),
3537            connection_max_retries: Some(1),
3538            reconnect_max_attempts: None,
3539            heartbeat_timeout_secs: None,
3540            certs_dir: None,
3541        };
3542
3543        // Should successfully connect with URL format
3544        let client = SocketClient::builder().config(config).connect().await;
3545        assert!(
3546            client.is_ok(),
3547            "Client should connect with URL scheme format"
3548        );
3549
3550        if let Ok(client) = client {
3551            client.close().await;
3552        }
3553        server.abort();
3554    }
3555
3556    #[rstest]
3557    fn test_parse_socket_url_ipv6_with_zone() {
3558        // IPv6 with zone ID (link-local address)
3559        let (socket_addr, request_url) =
3560            SocketClientInner::parse_socket_url("[fe80::1%eth0]:8080", Mode::Plain).unwrap();
3561        assert_eq!(socket_addr, "[fe80::1%eth0]:8080");
3562        assert_eq!(request_url, "ws://[fe80::1%eth0]:8080");
3563
3564        // Verify zone is preserved in URL format too
3565        let (socket_addr, request_url) =
3566            SocketClientInner::parse_socket_url("ws://[fe80::1%lo]:9090", Mode::Plain).unwrap();
3567        assert_eq!(socket_addr, "[fe80::1%lo]:9090");
3568        assert_eq!(request_url, "ws://[fe80::1%lo]:9090");
3569    }
3570
3571    #[rstest]
3572    #[tokio::test]
3573    async fn test_ipv6_loopback_connection() {
3574        // Test IPv6 loopback address connection
3575        // Skip if IPv6 is not available on the system
3576        if TcpListener::bind("[::1]:0").await.is_err() {
3577            return;
3578        }
3579
3580        let listener = TcpListener::bind("[::1]:0").await.unwrap();
3581        let port = listener.local_addr().unwrap().port();
3582
3583        let server = task::spawn(async move {
3584            if let Ok((mut sock, _)) = listener.accept().await {
3585                let mut buf = vec![0u8; 1024];
3586                if let Ok(n) = sock.read(&mut buf).await {
3587                    // Echo back
3588                    let _ = sock.write_all(&buf[..n]).await;
3589                }
3590            }
3591            sleep(Duration::from_millis(50)).await;
3592        });
3593
3594        let config = SocketConfig {
3595            url: format!("[::1]:{port}"), // IPv6 loopback
3596            mode: Mode::Plain,
3597            suffix: b"\r\n".to_vec(),
3598            message_handler: None,
3599            heartbeat: None,
3600            connect_timeout_ms: Some(1_000),
3601            reconnect_delay_initial_ms: Some(50),
3602            reconnect_delay_max_ms: Some(100),
3603            reconnect_backoff_factor: Some(1.0),
3604            reconnect_jitter_ms: Some(0),
3605            connection_max_retries: Some(1),
3606            reconnect_max_attempts: None,
3607            heartbeat_timeout_secs: None,
3608            certs_dir: None,
3609        };
3610
3611        let client = SocketClient::builder().config(config).connect().await;
3612        assert!(
3613            client.is_ok(),
3614            "Client should connect to IPv6 loopback address"
3615        );
3616
3617        if let Ok(client) = client {
3618            client.close().await;
3619        }
3620        server.abort();
3621    }
3622
3623    #[rstest]
3624    #[tokio::test]
3625    async fn test_send_waits_during_reconnection() {
3626        // Test that send operations wait for reconnection to complete (up to configured timeout)
3627        use nautilus_common::testing::wait_until_async;
3628
3629        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3630        let port = listener.local_addr().unwrap().port();
3631
3632        let server = task::spawn(async move {
3633            // First connection - accept and immediately close
3634            if let Ok((sock, _)) = listener.accept().await {
3635                drop(sock);
3636            }
3637
3638            // Wait before accepting second connection
3639            sleep(Duration::from_millis(500)).await;
3640
3641            // Second connection - accept and keep alive
3642            if let Ok((mut sock, _)) = listener.accept().await {
3643                // Echo messages
3644                let mut buf = vec![0u8; 1024];
3645                while let Ok(n) = sock.read(&mut buf).await {
3646                    if n == 0 {
3647                        break;
3648                    }
3649
3650                    if sock.write_all(&buf[..n]).await.is_err() {
3651                        break;
3652                    }
3653                }
3654            }
3655        });
3656
3657        let config = SocketConfig {
3658            url: format!("127.0.0.1:{port}"),
3659            mode: Mode::Plain,
3660            suffix: b"\r\n".to_vec(),
3661            message_handler: None,
3662            heartbeat: None,
3663            connect_timeout_ms: Some(5_000), // 5s timeout - enough for reconnect
3664            reconnect_delay_initial_ms: Some(100),
3665            reconnect_delay_max_ms: Some(200),
3666            reconnect_backoff_factor: Some(1.0),
3667            reconnect_jitter_ms: Some(0),
3668            connection_max_retries: Some(1),
3669            reconnect_max_attempts: None,
3670            heartbeat_timeout_secs: None,
3671            certs_dir: None,
3672        };
3673
3674        let client = SocketClient::builder()
3675            .config(config)
3676            .connect()
3677            .await
3678            .unwrap();
3679
3680        // Wait for reconnection to trigger
3681        wait_until_async(
3682            || async { client.is_reconnecting() },
3683            Duration::from_secs(2),
3684        )
3685        .await;
3686
3687        // Try to send while reconnecting - should wait and succeed after reconnect
3688        let send_result = tokio::time::timeout(
3689            Duration::from_secs(3),
3690            client.send_bytes(b"test_message".to_vec()),
3691        )
3692        .await;
3693
3694        assert!(
3695            send_result.is_ok() && send_result.unwrap().is_ok(),
3696            "Send should succeed after waiting for reconnection"
3697        );
3698
3699        client.close().await;
3700        server.abort();
3701    }
3702
3703    #[rstest]
3704    #[tokio::test]
3705    async fn test_send_bytes_timeout_uses_configured_connect_timeout() {
3706        // Test that send_bytes operations respect the configured connect_timeout.
3707        // When a client is stuck in RECONNECT longer than the timeout, sends should fail with Timeout.
3708        use nautilus_common::testing::wait_until_async;
3709
3710        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3711        let port = listener.local_addr().unwrap().port();
3712
3713        let server = task::spawn(async move {
3714            // Accept first connection and immediately close it
3715            if let Ok((sock, _)) = listener.accept().await {
3716                drop(sock);
3717            }
3718            // Drop listener entirely so reconnection fails completely
3719            drop(listener);
3720            sleep(Duration::from_mins(1)).await;
3721        });
3722
3723        let config = SocketConfig {
3724            url: format!("127.0.0.1:{port}"),
3725            mode: Mode::Plain,
3726            suffix: b"\r\n".to_vec(),
3727            message_handler: None,
3728            heartbeat: None,
3729            connect_timeout_ms: Some(1_000), // 1s timeout for faster test
3730            reconnect_delay_initial_ms: Some(200), // Short backoff (but > timeout) to keep client in RECONNECT
3731            reconnect_delay_max_ms: Some(200),
3732            reconnect_backoff_factor: Some(1.0),
3733            reconnect_jitter_ms: Some(0),
3734            connection_max_retries: Some(1),
3735            reconnect_max_attempts: None,
3736            heartbeat_timeout_secs: None,
3737            certs_dir: None,
3738        };
3739
3740        let client = SocketClient::builder()
3741            .config(config)
3742            .connect()
3743            .await
3744            .unwrap();
3745
3746        // Wait for client to enter RECONNECT state
3747        wait_until_async(
3748            || async { client.is_reconnecting() },
3749            Duration::from_secs(3),
3750        )
3751        .await;
3752
3753        // Attempt send while stuck in RECONNECT - should timeout after 1s (configured timeout)
3754        // The client will try to reconnect for 1s, fail, then wait 5s backoff before next attempt
3755        let start = std::time::Instant::now();
3756        let send_result = client.send_bytes(b"test".to_vec()).await;
3757        let elapsed = start.elapsed();
3758
3759        assert!(
3760            send_result.is_err(),
3761            "Send should fail when client stuck in RECONNECT, was: {send_result:?}"
3762        );
3763        assert!(
3764            matches!(send_result, Err(crate::error::SendError::Timeout)),
3765            "Send should return Timeout error, was: {send_result:?}"
3766        );
3767        // Verify timeout respects configured value (1s), but don't check upper bound
3768        // as CI scheduler jitter can cause legitimate delays beyond the timeout
3769        assert!(
3770            elapsed >= Duration::from_millis(900),
3771            "Send should timeout after at least 1s (configured timeout), took {elapsed:?}"
3772        );
3773
3774        client.close().await;
3775        server.abort();
3776    }
3777
3778    #[rstest]
3779    #[tokio::test]
3780    async fn test_heartbeat_timeout_triggers_reconnect() {
3781        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3782        let port = listener.local_addr().unwrap().port();
3783
3784        // Server accepts connection but sends nothing (simulates silent death)
3785        let server = task::spawn(async move {
3786            let (_sock1, _) = listener.accept().await.unwrap();
3787            // Hold connection open but send nothing, wait for reconnect attempt
3788            sleep(Duration::from_secs(5)).await;
3789        });
3790
3791        let config = SocketConfig {
3792            url: format!("127.0.0.1:{port}"),
3793            mode: Mode::Plain,
3794            suffix: b"\r\n".to_vec(),
3795            message_handler: None,
3796            heartbeat: None,
3797            connect_timeout_ms: Some(2_000),
3798            reconnect_delay_initial_ms: Some(50),
3799            reconnect_delay_max_ms: Some(100),
3800            reconnect_backoff_factor: Some(1.0),
3801            reconnect_jitter_ms: Some(0),
3802            connection_max_retries: Some(1),
3803            reconnect_max_attempts: Some(1),
3804            heartbeat_timeout_secs: Some(1),
3805            certs_dir: None,
3806        };
3807
3808        let client = SocketClient::builder()
3809            .config(config)
3810            .connect()
3811            .await
3812            .unwrap();
3813
3814        assert!(client.is_active());
3815
3816        // Wait for the dead-peer timeout to fire and the client to enter reconnect
3817        wait_until_async(
3818            || async { client.is_reconnecting() || client.is_closed() },
3819            Duration::from_secs(4),
3820        )
3821        .await;
3822
3823        assert!(
3824            !client.is_active(),
3825            "Client should not be active after the dead-peer timeout"
3826        );
3827
3828        client.close().await;
3829        server.abort();
3830    }
3831
3832    #[rstest]
3833    #[tokio::test]
3834    async fn test_heartbeat_timeout_resets_on_inbound_bytes() {
3835        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3836        let port = listener.local_addr().unwrap().port();
3837
3838        // Server sends data every 200ms (well within the 1s idle timeout)
3839        let server = task::spawn(async move {
3840            let (mut sock, _) = listener.accept().await.unwrap();
3841            for _ in 0..10 {
3842                sleep(Duration::from_millis(200)).await;
3843
3844                if sock.write_all(b"ping\r\n").await.is_err() {
3845                    break;
3846                }
3847            }
3848        });
3849
3850        let config = SocketConfig {
3851            url: format!("127.0.0.1:{port}"),
3852            mode: Mode::Plain,
3853            suffix: b"\r\n".to_vec(),
3854            message_handler: None,
3855            heartbeat: None,
3856            connect_timeout_ms: Some(2_000),
3857            reconnect_delay_initial_ms: Some(50),
3858            reconnect_delay_max_ms: Some(100),
3859            reconnect_backoff_factor: Some(1.0),
3860            reconnect_jitter_ms: Some(0),
3861            connection_max_retries: Some(1),
3862            reconnect_max_attempts: Some(1),
3863            heartbeat_timeout_secs: Some(1),
3864            certs_dir: None,
3865        };
3866
3867        let client = SocketClient::builder()
3868            .config(config)
3869            .connect()
3870            .await
3871            .unwrap();
3872
3873        assert!(client.is_active());
3874
3875        // Wait 1.5s - data arrives every 200ms so idle timeout (1s) should NOT fire
3876        sleep(Duration::from_millis(1_500)).await;
3877
3878        assert!(
3879            client.is_active(),
3880            "Client should remain active when data is flowing"
3881        );
3882
3883        client.close().await;
3884        server.abort();
3885    }
3886
3887    #[rstest]
3888    #[tokio::test]
3889    async fn test_close_during_backoff_exits_promptly() {
3890        // Verify that close() interrupts backoff sleep (Finding 1).
3891        // Server accepts then drops, no second listener -> reconnect fails -> enters backoff.
3892        // We close while backing off and assert the client shuts down quickly.
3893        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3894        let port = listener.local_addr().unwrap().port();
3895
3896        let server = task::spawn(async move {
3897            // Accept first connection, close immediately
3898            if let Ok((mut sock, _)) = listener.accept().await {
3899                drop(sock.shutdown());
3900            }
3901            // Don't accept again so reconnect fails and enters backoff
3902            sleep(Duration::from_mins(1)).await;
3903        });
3904
3905        let config = SocketConfig {
3906            url: format!("127.0.0.1:{port}"),
3907            mode: Mode::Plain,
3908            suffix: b"\r\n".to_vec(),
3909            message_handler: None,
3910            heartbeat: None,
3911            connect_timeout_ms: Some(1_000),
3912            reconnect_delay_initial_ms: Some(10_000), // 10s backoff to ensure we're sleeping
3913            reconnect_delay_max_ms: Some(10_000),
3914            reconnect_backoff_factor: Some(1.0),
3915            reconnect_jitter_ms: Some(0),
3916            connection_max_retries: None,
3917            reconnect_max_attempts: None,
3918            heartbeat_timeout_secs: None,
3919            certs_dir: None,
3920        };
3921
3922        let client = SocketClient::builder()
3923            .config(config)
3924            .connect()
3925            .await
3926            .unwrap();
3927
3928        // Wait for client to enter reconnect
3929        wait_until_async(
3930            || async { client.is_reconnecting() },
3931            Duration::from_secs(3),
3932        )
3933        .await;
3934
3935        // Wait for the reconnect attempt to fail and enter backoff sleep
3936        sleep(Duration::from_millis(1_500)).await;
3937
3938        // Close while backing off
3939        let start = std::time::Instant::now();
3940        client.close().await;
3941        let elapsed = start.elapsed();
3942
3943        assert!(client.is_closed(), "Client should be closed");
3944        // Should exit well before the 10s backoff sleep completes
3945        assert!(
3946            elapsed < Duration::from_secs(2),
3947            "Close should interrupt backoff sleep, took {elapsed:?}"
3948        );
3949
3950        server.abort();
3951    }
3952
3953    #[rstest]
3954    #[tokio::test]
3955    async fn test_zero_heartbeat_timeout_rejected() {
3956        let config = SocketConfig {
3957            url: "127.0.0.1:9999".to_string(),
3958            mode: Mode::Plain,
3959            suffix: b"\r\n".to_vec(),
3960            message_handler: None,
3961            heartbeat: None,
3962            connect_timeout_ms: None,
3963            reconnect_delay_initial_ms: None,
3964            reconnect_delay_max_ms: None,
3965            reconnect_backoff_factor: None,
3966            reconnect_jitter_ms: None,
3967            reconnect_max_attempts: None,
3968            connection_max_retries: Some(1),
3969            heartbeat_timeout_secs: Some(0),
3970            certs_dir: None,
3971        };
3972
3973        let result = SocketClient::builder().config(config).connect().await;
3974
3975        assert!(result.is_err(), "Zero heartbeat timeout should be rejected");
3976        let err_msg = result.unwrap_err().to_string();
3977        assert!(
3978            err_msg.contains("heartbeat_timeout_secs"),
3979            "Error should name the offending field, was: {err_msg}"
3980        );
3981    }
3982
3983    #[rstest]
3984    #[tokio::test]
3985    async fn test_empty_suffix_rejected() {
3986        let config = SocketConfig {
3987            url: "127.0.0.1:9999".to_string(),
3988            mode: Mode::Plain,
3989            suffix: vec![],
3990            message_handler: None,
3991            heartbeat: None,
3992            connect_timeout_ms: None,
3993            reconnect_delay_initial_ms: None,
3994            reconnect_delay_max_ms: None,
3995            reconnect_backoff_factor: None,
3996            reconnect_jitter_ms: None,
3997            reconnect_max_attempts: None,
3998            connection_max_retries: Some(1),
3999            heartbeat_timeout_secs: None,
4000            certs_dir: None,
4001        };
4002
4003        let result = SocketClient::builder().config(config).connect().await;
4004
4005        assert!(
4006            result.is_err(),
4007            "Empty suffix should cause connection to fail"
4008        );
4009        let err_msg = result.unwrap_err().to_string();
4010        assert!(
4011            err_msg.contains("suffix cannot be empty"),
4012            "Error should mention empty suffix, was: {err_msg}"
4013        );
4014    }
4015}
4016
4017#[cfg(test)]
4018mod reconnect_request_tests {
4019    use std::sync::{Arc, atomic::AtomicU8};
4020
4021    use parking_lot::Mutex;
4022    use rstest::rstest;
4023
4024    use super::*;
4025    use crate::SocketState;
4026
4027    fn handle(
4028        mode: ConnectionMode,
4029    ) -> (
4030        SocketReconnectHandle,
4031        Arc<tokio::sync::Notify>,
4032        Arc<Mutex<Vec<SocketState>>>,
4033    ) {
4034        let controller_notify = Arc::new(tokio::sync::Notify::new());
4035        let states = Arc::new(Mutex::new(Vec::new()));
4036        let states_callback = Arc::clone(&states);
4037        let state_sink = SocketStateSink::new(move |state| {
4038            states_callback.lock().push(state);
4039        });
4040        let handle = SocketReconnectHandle {
4041            connection_mode: Arc::new(AtomicU8::new(mode.as_u8())),
4042            state_sink: Some(state_sink),
4043            controller_lifecycle: Arc::new(ControllerLifecycle::new()),
4044            controller_notify: Arc::clone(&controller_notify),
4045        };
4046        (handle, controller_notify, states)
4047    }
4048
4049    #[rstest]
4050    #[tokio::test]
4051    async fn accepted_request_reports_loss_and_wakes_controller_once() {
4052        let (handle, controller_notify, states) = handle(ConnectionMode::Active);
4053
4054        assert_eq!(
4055            handle.request_reconnect(),
4056            ReconnectRequestOutcome::Accepted
4057        );
4058        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
4059        tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4060            .await
4061            .expect("accepted request should notify controller");
4062
4063        assert_eq!(
4064            handle.request_reconnect(),
4065            ReconnectRequestOutcome::AlreadyReconnecting
4066        );
4067        assert_eq!(*states.lock(), vec![SocketState::Disconnected]);
4068        assert!(
4069            tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4070                .await
4071                .is_err(),
4072            "duplicate request should not notify controller",
4073        );
4074    }
4075
4076    #[rstest]
4077    #[case(
4078        ConnectionMode::Reconnect,
4079        ReconnectRequestOutcome::AlreadyReconnecting
4080    )]
4081    #[case(ConnectionMode::Disconnect, ReconnectRequestOutcome::Disconnected)]
4082    #[case(ConnectionMode::Closed, ReconnectRequestOutcome::Closed)]
4083    #[tokio::test]
4084    async fn rejected_request_preserves_state_and_does_not_wake_controller(
4085        #[case] mode: ConnectionMode,
4086        #[case] expected: ReconnectRequestOutcome,
4087    ) {
4088        let (handle, controller_notify, states) = handle(mode);
4089
4090        assert_eq!(handle.request_reconnect(), expected);
4091        assert!(states.lock().is_empty());
4092        assert!(
4093            tokio::time::timeout(Duration::from_millis(10), controller_notify.notified())
4094                .await
4095                .is_err(),
4096            "rejected request should not notify controller",
4097        );
4098    }
4099}