Skip to main content

nntp_proxy/pool/
deadpool_connection.rs

1//! Core TCP connection manager for deadpool
2//!
3//! This module provides the `TcpManager` struct which handles the low-level
4//! creation of optimized TCP/TLS connections to NNTP servers.
5
6use deadpool::managed;
7use std::collections::VecDeque;
8use std::io;
9use std::net::{IpAddr, SocketAddr};
10use std::sync::Arc;
11use std::sync::atomic::{AtomicUsize, Ordering};
12use tokio::io::AsyncWriteExt;
13use tokio::net::TcpStream;
14use tokio::sync::Mutex;
15use tokio::sync::Notify;
16use tokio::sync::RwLock;
17
18use crate::connection_error::ConnectionError;
19use crate::protocol::{RequestContext, authinfo_pass, authinfo_user};
20use crate::stream::ConnectionStream;
21use crate::tls::{TlsConfig, TlsManager};
22
23/// Type alias for the deadpool connection pool
24pub type Pool = managed::Pool<TcpManager>;
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27enum CompressionSupport {
28    Supported,
29    Unsupported,
30}
31
32#[derive(Debug, Default)]
33enum CompressionSupportState {
34    #[default]
35    Unknown,
36    Probing(Arc<Notify>),
37    Supported,
38    Unsupported,
39}
40
41/// Optional settings for [`TcpManager`] construction
42///
43/// Groups optional parameters (credentials, TLS, compression) to keep
44/// the `TcpManager::new()` signature concise.
45#[derive(Debug, Clone)]
46pub struct TcpManagerOptions {
47    pub username: Option<String>,
48    pub password: Option<String>,
49    pub tls_config: Option<TlsConfig>,
50    /// TCP receive buffer size for this connection.
51    pub recv_buffer_size: usize,
52    /// TCP send buffer size for this connection.
53    pub send_buffer_size: usize,
54    /// Wire compression mode: None = auto-detect, Some(true) = require, Some(false) = disable
55    pub compress: Option<bool>,
56    /// Compression level (0-9). None = fast (level 1).
57    pub compress_level: Option<u32>,
58    /// Send MODE READER to the backend after authentication.
59    ///
60    /// RFC 3977 §5.3: MODE READER switches a transit server to reading mode.
61    /// Defaults to `true` — required for reader-capable backends.
62    pub send_mode_reader: bool,
63}
64
65impl Default for TcpManagerOptions {
66    fn default() -> Self {
67        Self {
68            username: None,
69            password: None,
70            tls_config: None,
71            recv_buffer_size: crate::constants::socket::HIGH_THROUGHPUT_RECV_BUFFER,
72            send_buffer_size: crate::constants::socket::HIGH_THROUGHPUT_SEND_BUFFER,
73            compress: None,
74            compress_level: None,
75            send_mode_reader: true,
76        }
77    }
78}
79
80/// TCP connection manager for deadpool with cached TLS config
81#[derive(Debug, Clone)]
82pub struct TcpManager {
83    pub(crate) host: String,
84    pub(crate) port: u16,
85    pub(crate) name: String,
86    pub(crate) username: Option<String>,
87    pub(crate) password: Option<String>,
88    pub(crate) tls_config: TlsConfig,
89    /// Cached TLS manager with pre-loaded certificates (avoids base64 decode overhead)
90    pub(crate) tls_manager: Option<Arc<TlsManager>>,
91    resolved_socket_addrs: Arc<RwLock<Option<Arc<[SocketAddr]>>>>,
92    next_resolved_socket_addr: Arc<AtomicUsize>,
93    pub(crate) recv_buffer_size: usize,
94    pub(crate) send_buffer_size: usize,
95    /// Wire compression mode: None = auto-detect, Some(true) = require, Some(false) = disable
96    pub(crate) compress: Option<bool>,
97    /// Compression level (0-9). None = fast (level 1).
98    pub(crate) compress_level: Option<u32>,
99    compression_support: Arc<Mutex<CompressionSupportState>>,
100    /// Whether to send MODE READER after authentication (RFC 3977 §5.3)
101    pub(crate) send_mode_reader: bool,
102}
103
104impl TcpManager {
105    fn socket_buffer_size_u32(size: usize, label: &str) -> Result<u32, ConnectionError> {
106        u32::try_from(size).map_err(|_| {
107            ConnectionError::IoError(io::Error::new(
108                io::ErrorKind::InvalidInput,
109                format!("{label} socket buffer size {size} exceeds u32::MAX"),
110            ))
111        })
112    }
113
114    fn ip_literal_socket_addr(&self) -> Option<SocketAddr> {
115        self.host
116            .parse::<IpAddr>()
117            .ok()
118            .map(|ip| SocketAddr::new(ip, self.port))
119    }
120
121    async fn resolve_socket_addrs(&self) -> Result<Arc<[SocketAddr]>, ConnectionError> {
122        if let Some(socket_addr) = self.ip_literal_socket_addr() {
123            return Ok(Arc::from([socket_addr]));
124        }
125
126        if let Some(addrs) = self.resolved_socket_addrs.read().await.as_ref() {
127            if addrs.is_empty() {
128                return Err(ConnectionError::DnsNoAddresses {
129                    address: format!("{}:{}", self.host, self.port),
130                });
131            }
132            return Ok(addrs.clone());
133        }
134
135        let mut cached_addrs = self.resolved_socket_addrs.write().await;
136        if let Some(addrs) = cached_addrs.as_ref() {
137            if addrs.is_empty() {
138                return Err(ConnectionError::DnsNoAddresses {
139                    address: format!("{}:{}", self.host, self.port),
140                });
141            }
142            return Ok(addrs.clone());
143        }
144
145        let addrs = tokio::net::lookup_host((self.host.as_str(), self.port))
146            .await?
147            .collect::<Vec<_>>();
148        if addrs.is_empty() {
149            return Err(ConnectionError::DnsNoAddresses {
150                address: format!("{}:{}", self.host, self.port),
151            });
152        }
153
154        let addrs: Arc<[SocketAddr]> = Arc::from(addrs);
155        *cached_addrs = Some(addrs.clone());
156        Ok(addrs)
157    }
158
159    async fn refresh_socket_addrs(&self) -> Result<Arc<[SocketAddr]>, ConnectionError> {
160        if let Some(socket_addr) = self.ip_literal_socket_addr() {
161            return Ok(Arc::from([socket_addr]));
162        }
163
164        let addrs = tokio::net::lookup_host((self.host.as_str(), self.port))
165            .await?
166            .collect::<Vec<_>>();
167        if addrs.is_empty() {
168            return Err(ConnectionError::DnsNoAddresses {
169                address: format!("{}:{}", self.host, self.port),
170            });
171        }
172
173        let addrs: Arc<[SocketAddr]> = Arc::from(addrs);
174        *self.resolved_socket_addrs.write().await = Some(addrs.clone());
175        Ok(addrs)
176    }
177
178    fn is_ipv6_network_unreachable(socket_addr: SocketAddr, error: &ConnectionError) -> bool {
179        socket_addr.is_ipv6()
180            && matches!(
181                error,
182                ConnectionError::IoError(error)
183                    if matches!(
184                        error.kind(),
185                        io::ErrorKind::NetworkUnreachable | io::ErrorKind::HostUnreachable
186                    )
187            )
188    }
189
190    async fn remove_cached_ipv6_socket_addrs(&self) {
191        let mut cached_addrs = self.resolved_socket_addrs.write().await;
192        let Some(addrs) = cached_addrs.as_ref() else {
193            return;
194        };
195
196        let ipv4_addrs = addrs
197            .iter()
198            .copied()
199            .filter(SocketAddr::is_ipv4)
200            .collect::<Vec<_>>();
201        if ipv4_addrs.len() == addrs.len() {
202            return;
203        }
204
205        *cached_addrs = if ipv4_addrs.is_empty() {
206            None
207        } else {
208            Some(Arc::<[SocketAddr]>::from(ipv4_addrs))
209        };
210    }
211
212    /// Create a new `TcpManager` with optional TLS configuration
213    ///
214    /// If `options.tls_config` is `Some` with `use_tls = true`, the TLS manager is
215    /// pre-initialized (certificates loaded). If `None` or `use_tls = false`,
216    /// plain TCP connections are used.
217    ///
218    /// # Errors
219    /// Returns any TLS initialization error when TLS is enabled for the manager.
220    pub fn new(
221        host: String,
222        port: u16,
223        name: String,
224        options: TcpManagerOptions,
225    ) -> Result<Self, ConnectionError> {
226        let (tls_config, tls_manager) = match options.tls_config {
227            Some(cfg) if cfg.use_tls => {
228                let mgr = Arc::new(TlsManager::new(cfg.clone()).map_err(|e| {
229                    ConnectionError::TlsHandshake {
230                        backend: name.clone(),
231                        source: e.into(),
232                    }
233                })?);
234                (cfg, Some(mgr))
235            }
236            Some(cfg) => (cfg, None),
237            None => (TlsConfig::default(), None),
238        };
239
240        Ok(Self {
241            host,
242            port,
243            name,
244            username: options.username,
245            password: options.password,
246            tls_config,
247            tls_manager,
248            resolved_socket_addrs: Arc::new(RwLock::new(None)),
249            next_resolved_socket_addr: Arc::new(AtomicUsize::new(0)),
250            recv_buffer_size: options.recv_buffer_size,
251            send_buffer_size: options.send_buffer_size,
252            compress: options.compress,
253            compress_level: options.compress_level,
254            compression_support: Arc::new(Mutex::new(CompressionSupportState::Unknown)),
255            send_mode_reader: options.send_mode_reader,
256        })
257    }
258
259    async fn connect_socket_addr(
260        &self,
261        socket_addr: SocketAddr,
262    ) -> Result<TcpStream, ConnectionError> {
263        // Create tokio TcpSocket (non-blocking from the start)
264        let socket = if socket_addr.is_ipv4() {
265            tokio::net::TcpSocket::new_v4()?
266        } else {
267            tokio::net::TcpSocket::new_v6()?
268        };
269
270        // Pre-connect options (buffer sizes, reuse)
271        if self.recv_buffer_size > 0 {
272            socket.set_recv_buffer_size(Self::socket_buffer_size_u32(
273                self.recv_buffer_size,
274                "receive",
275            )?)?;
276        }
277        if self.send_buffer_size > 0 {
278            socket.set_send_buffer_size(Self::socket_buffer_size_u32(
279                self.send_buffer_size,
280                "send",
281            )?)?;
282        }
283        socket.set_reuseaddr(true)?;
284
285        // Async connect — does NOT block the tokio worker thread
286        let tcp_stream = socket.connect(socket_addr).await?;
287
288        // Post-connect options via socket2::SockRef (keepalive with params, nodelay)
289        let sock_ref = socket2::SockRef::from(&tcp_stream);
290        sock_ref.set_keepalive(true)?;
291        let keepalive = socket2::TcpKeepalive::new()
292            .with_time(crate::constants::duration_polyfill::from_minutes(1))
293            .with_interval(std::time::Duration::from_secs(10));
294        sock_ref.set_tcp_keepalive(&keepalive)?;
295        sock_ref.set_tcp_nodelay(true)?;
296
297        Ok(tcp_stream)
298    }
299
300    async fn create_connected_tcp_stream(&self) -> Result<TcpStream, ConnectionError> {
301        let addrs = self.resolve_socket_addrs().await?;
302        let last_error = match self.try_resolved_socket_addrs(&addrs).await {
303            Ok(tcp_stream) => return Ok(tcp_stream),
304            Err(last_error) => last_error,
305        };
306
307        if self.ip_literal_socket_addr().is_some() {
308            return Err(
309                last_error.unwrap_or_else(|| ConnectionError::DnsNoAddresses {
310                    address: format!("{}:{}", self.host, self.port),
311                }),
312            );
313        }
314
315        tracing::debug!(
316            backend = %self.name,
317            host = %self.host,
318            "All cached backend socket addresses failed; refreshing DNS before final connect pass"
319        );
320
321        let addrs = self.refresh_socket_addrs().await?;
322        self.try_resolved_socket_addrs(&addrs)
323            .await
324            .map_err(|last_error| {
325                last_error.unwrap_or_else(|| ConnectionError::DnsNoAddresses {
326                    address: format!("{}:{}", self.host, self.port),
327                })
328            })
329    }
330
331    async fn try_resolved_socket_addrs(
332        &self,
333        addrs: &[SocketAddr],
334    ) -> Result<TcpStream, Option<ConnectionError>> {
335        let start = self
336            .next_resolved_socket_addr
337            .fetch_add(1, Ordering::Relaxed)
338            % addrs.len();
339        let mut last_error = None;
340
341        let mut remaining_addrs = (0..addrs.len())
342            .map(|offset| addrs[(start + offset) % addrs.len()])
343            .collect::<VecDeque<_>>();
344
345        while let Some(socket_addr) = remaining_addrs.pop_front() {
346            match self.connect_socket_addr(socket_addr).await {
347                Ok(tcp_stream) => return Ok(tcp_stream),
348                Err(error) => {
349                    if Self::is_ipv6_network_unreachable(socket_addr, &error) {
350                        self.remove_cached_ipv6_socket_addrs().await;
351                        remaining_addrs.retain(SocketAddr::is_ipv4);
352                    }
353
354                    tracing::debug!(
355                        backend = %self.name,
356                        host = %self.host,
357                        socket_addr = %socket_addr,
358                        error = %error,
359                        "Backend socket address connect failed; trying next resolved address"
360                    );
361                    last_error = Some(error);
362                }
363            }
364        }
365
366        Err(last_error)
367    }
368
369    /// Create an optimized connection (TCP or TLS)
370    pub(crate) async fn create_optimized_stream(
371        &self,
372    ) -> Result<ConnectionStream, ConnectionError> {
373        let tcp_stream = self.create_connected_tcp_stream().await?;
374
375        // Perform TLS handshake if enabled
376        if self.tls_config.use_tls {
377            // Use cached TLS manager to avoid re-parsing certificates
378            let Some(tls_manager) = self.tls_manager.as_ref() else {
379                return Err(ConnectionError::TlsHandshake {
380                    backend: self.name.clone(),
381                    source: "TLS enabled but TLS manager not initialized".into(),
382                });
383            };
384
385            let tls_stream = tls_manager
386                .handshake(tcp_stream, &self.host, &self.name)
387                .await
388                .map_err(|e| ConnectionError::TlsHandshake {
389                    backend: self.name.clone(),
390                    source: e.into(),
391                })?;
392            Ok(ConnectionStream::tls(tls_stream))
393        } else {
394            Ok(ConnectionStream::plain(tcp_stream))
395        }
396    }
397}
398
399// ============================================================================
400// Connection setup: greeting, auth, and future negotiation hooks
401// ============================================================================
402
403impl TcpManager {
404    /// Read one single-line setup reply using the same backend reply framing
405    /// facade as normal commands.
406    ///
407    /// Connection setup commands are not on the hot article path, but they must
408    /// still tolerate split TCP reads and must not open another place that
409    /// reasons about response line boundaries.
410    async fn read_backend_setup_reply(
411        stream: &mut ConnectionStream,
412        request: &RequestContext,
413        buffer: &mut [u8],
414    ) -> Result<String, ConnectionError> {
415        match crate::session::backend::read_single_line_reply(stream, request, buffer).await {
416            Ok(reply) => Ok(reply),
417            Err(
418                crate::session::backend::SingleLineReplyReadError::Full { bytes_read }
419                | crate::session::backend::SingleLineReplyReadError::Invalid { bytes_read },
420            ) => Err(ConnectionError::IoError(std::io::Error::new(
421                std::io::ErrorKind::InvalidData,
422                format!(
423                    "invalid or truncated backend setup reply: {}",
424                    String::from_utf8_lossy(&buffer[..bytes_read]).trim_end()
425                ),
426            ))),
427            Err(crate::session::backend::SingleLineReplyReadError::Io(err)) => Err(err.into()),
428            Err(crate::session::backend::SingleLineReplyReadError::Closed) => {
429                Err(ConnectionError::IoError(std::io::Error::new(
430                    std::io::ErrorKind::UnexpectedEof,
431                    "backend closed while reading setup reply",
432                )))
433            }
434        }
435    }
436
437    /// Read and validate the NNTP server greeting.
438    async fn consume_greeting(
439        &self,
440        stream: &mut ConnectionStream,
441        buffer: &mut [u8],
442    ) -> Result<(), ConnectionError> {
443        let request = RequestContext::from_verb_args(b"MODE", b"READER");
444        let greeting = Self::read_backend_setup_reply(stream, &request, buffer).await?;
445
446        if !crate::protocol::StatusCode::parse(greeting.as_bytes())
447            .is_some_and(|code| code.is_greeting())
448        {
449            return Err(ConnectionError::InvalidGreeting {
450                backend: self.name.clone(),
451                greeting: greeting.trim().to_string(),
452            });
453        }
454
455        Ok(())
456    }
457
458    /// Send MODE READER to the backend after authentication.
459    ///
460    /// RFC 3977 §5.3: MODE READER switches the server into reader mode.
461    /// Valid responses are 200 (posting allowed) or 201 (posting not permitted).
462    /// Any other response indicates the service is unavailable for reading.
463    async fn negotiate_mode_reader(
464        &self,
465        stream: &mut ConnectionStream,
466        buffer: &mut [u8],
467    ) -> Result<(), ConnectionError> {
468        stream.write_all(b"MODE READER\r\n").await?;
469        stream.flush().await?;
470
471        let request = RequestContext::from_verb_args(b"MODE", b"READER");
472        let response = Self::read_backend_setup_reply(stream, &request, buffer).await?;
473
474        // RFC 3977 §5.3: 200 = reader mode + posting allowed, 201 = reader mode + posting not permitted
475        if crate::protocol::StatusCode::parse(response.as_bytes())
476            .is_some_and(|code| matches!(code.as_u16(), 200 | 201))
477        {
478            tracing::debug!(
479                backend = %self.name,
480                response = %response.trim(),
481                "MODE READER accepted"
482            );
483            return Ok(());
484        }
485
486        Err(ConnectionError::InvalidGreeting {
487            backend: self.name.clone(),
488            greeting: response.trim().to_string(),
489        })
490    }
491
492    /// Negotiate COMPRESS DEFLATE (RFC 8054) with the backend server.
493    ///
494    /// Returns `Ok(true)` if compression was successfully negotiated,
495    /// `Ok(false)` if compression was skipped or not supported.
496    async fn negotiate_compression(
497        &self,
498        stream: &mut ConnectionStream,
499        buffer: &mut [u8],
500    ) -> Result<bool, ConnectionError> {
501        if self.compress == Some(false) {
502            return Ok(false);
503        }
504
505        if self.compress == Some(true) {
506            return self
507                .probe_compression_with_timeout(stream, buffer)
508                .await
509                .map(|support| matches!(support, CompressionSupport::Supported));
510        }
511
512        loop {
513            let probe_waiter = {
514                let mut cached_support = self.compression_support.lock().await;
515                match &*cached_support {
516                    CompressionSupportState::Unsupported => {
517                        tracing::debug!(
518                            backend = %self.name,
519                            "Skipping COMPRESS DEFLATE; backend previously reported it unsupported"
520                        );
521                        return Ok(false);
522                    }
523                    CompressionSupportState::Supported => {
524                        drop(cached_support);
525                        return self
526                            .probe_compression_with_timeout(stream, buffer)
527                            .await
528                            .map(|support| matches!(support, CompressionSupport::Supported));
529                    }
530                    CompressionSupportState::Probing(notify) => {
531                        Some(notify.clone().notified_owned())
532                    }
533                    CompressionSupportState::Unknown => {
534                        let notify = Arc::new(Notify::new());
535                        *cached_support = CompressionSupportState::Probing(notify);
536                        None
537                    }
538                }
539            };
540
541            if let Some(waiter) = probe_waiter {
542                waiter.await;
543                continue;
544            }
545
546            let support = self.probe_compression_with_timeout(stream, buffer).await;
547            let mut cached_support = self.compression_support.lock().await;
548            let notify = match std::mem::take(&mut *cached_support) {
549                CompressionSupportState::Probing(notify) => notify,
550                state => {
551                    *cached_support = state;
552                    return support.map(|support| matches!(support, CompressionSupport::Supported));
553                }
554            };
555
556            match support {
557                Ok(CompressionSupport::Supported) => {
558                    *cached_support = CompressionSupportState::Supported;
559                    notify.notify_waiters();
560                    return Ok(true);
561                }
562                Ok(CompressionSupport::Unsupported) => {
563                    *cached_support = CompressionSupportState::Unsupported;
564                    notify.notify_waiters();
565                    return Ok(false);
566                }
567                Err(err) => {
568                    *cached_support = CompressionSupportState::Unknown;
569                    notify.notify_waiters();
570                    return Err(err);
571                }
572            }
573        }
574    }
575
576    async fn probe_compression_with_timeout(
577        &self,
578        stream: &mut ConnectionStream,
579        buffer: &mut [u8],
580    ) -> Result<CompressionSupport, ConnectionError> {
581        tokio::time::timeout(
582            crate::constants::timeout::CONNECTION,
583            self.probe_compression(stream, buffer),
584        )
585        .await
586        .map_err(|_| {
587            ConnectionError::IoError(io::Error::new(
588                io::ErrorKind::TimedOut,
589                "timed out negotiating COMPRESS DEFLATE",
590            ))
591        })?
592    }
593
594    async fn probe_compression(
595        &self,
596        stream: &mut ConnectionStream,
597        buffer: &mut [u8],
598    ) -> Result<CompressionSupport, ConnectionError> {
599        stream.write_all(crate::protocol::COMPRESS_DEFLATE).await?;
600        stream.flush().await?;
601
602        let request = RequestContext::from_verb_args(b"COMPRESS", b"DEFLATE");
603        let response = Self::read_backend_setup_reply(stream, &request, buffer).await?;
604
605        // 206 = Compression active
606        if crate::protocol::StatusCode::parse(response.as_bytes())
607            .is_some_and(|code| code.as_u16() == 206)
608        {
609            tracing::debug!(
610                backend = %self.name,
611                "COMPRESS DEFLATE negotiated successfully"
612            );
613            return Ok(CompressionSupport::Supported);
614        }
615
616        if self.compress == Some(true) {
617            return Err(ConnectionError::CompressionRequired {
618                backend: self.name.clone(),
619                response: response.trim().to_string(),
620            });
621        }
622
623        // Auto mode: compression not supported, continue without it
624        tracing::debug!(
625            backend = %self.name,
626            response = %response.trim(),
627            "COMPRESS DEFLATE not supported, continuing without compression"
628        );
629        Ok(CompressionSupport::Unsupported)
630    }
631
632    /// Perform AUTHINFO USER/PASS handshake if credentials are configured.
633    async fn negotiate_auth(
634        &self,
635        stream: &mut ConnectionStream,
636        buffer: &mut [u8],
637    ) -> Result<(), ConnectionError> {
638        let Some(username) = &self.username else {
639            return Ok(());
640        };
641
642        authinfo_user(username).write_wire_to(stream).await?;
643        let user_request = authinfo_user(username);
644        let response = Self::read_backend_setup_reply(stream, &user_request, buffer).await?;
645
646        if crate::protocol::StatusCode::parse(response.as_bytes())
647            .is_some_and(|code| code.requires_auth_credentials())
648        {
649            // Password required
650            let Some(password) = self.password.as_ref() else {
651                return Err(ConnectionError::PasswordRequired {
652                    backend: self.name.clone(),
653                });
654            };
655
656            authinfo_pass(password).write_wire_to(stream).await?;
657            let pass_request = authinfo_pass(password);
658            let response = Self::read_backend_setup_reply(stream, &pass_request, buffer).await?;
659
660            if !crate::protocol::StatusCode::parse(response.as_bytes())
661                .is_some_and(|code| code.is_auth_accepted())
662            {
663                // Check for 482 (connection limit exceeded) before generic auth failure
664                if crate::protocol::StatusCode::parse(response.as_bytes())
665                    .is_some_and(|c| c.as_u16() == 482)
666                {
667                    tracing::error!(
668                        backend = %self.name,
669                        host = %self.host,
670                        port = self.port,
671                        response = %response.trim(),
672                        "Backend connection limit exceeded"
673                    );
674                    return Err(ConnectionError::ConnectionLimitExceeded {
675                        backend: self.name.clone(),
676                        response: response.trim().to_string(),
677                    });
678                }
679
680                tracing::error!(
681                    "Authentication failed for {} ({}:{}) - Server response: {} - Username: {}",
682                    self.name,
683                    self.host,
684                    self.port,
685                    response.trim(),
686                    username
687                );
688                return Err(ConnectionError::AuthenticationFailed {
689                    backend: self.name.clone(),
690                    response: response.trim().to_string(),
691                });
692            }
693            tracing::debug!(
694                "Successfully authenticated to {} ({}:{}) as {}",
695                self.name,
696                self.host,
697                self.port,
698                username
699            );
700        } else if !crate::protocol::StatusCode::parse(response.as_bytes())
701            .is_some_and(|code| code.is_auth_accepted())
702        {
703            return Err(ConnectionError::UnexpectedAuthResponse {
704                backend: self.name.clone(),
705                response: response.trim().to_string(),
706            });
707        }
708
709        Ok(())
710    }
711}
712
713// ============================================================================
714// Manager trait implementation
715// ============================================================================
716
717impl managed::Manager for TcpManager {
718    type Type = ConnectionStream;
719    type Error = ConnectionError;
720
721    async fn create(&self) -> Result<ConnectionStream, ConnectionError> {
722        let mut stream = self.create_optimized_stream().await?;
723        let mut buffer = [0u8; 4096];
724
725        self.consume_greeting(&mut stream, &mut buffer).await?;
726        self.negotiate_auth(&mut stream, &mut buffer).await?;
727
728        if self.send_mode_reader {
729            self.negotiate_mode_reader(&mut stream, &mut buffer).await?;
730        }
731
732        if self.negotiate_compression(&mut stream, &mut buffer).await? {
733            let level = self.compress_level.unwrap_or(1);
734            stream = stream.into_compressed(level)?;
735        }
736
737        Ok(stream)
738    }
739
740    async fn recycle(
741        &self,
742        conn: &mut ConnectionStream,
743        _metrics: &managed::Metrics,
744    ) -> managed::RecycleResult<ConnectionError> {
745        use super::health_check::check_tcp_alive;
746        match check_tcp_alive(conn) {
747            Ok(()) => Ok(()),
748            Err(e) => {
749                // Shut down TCP immediately so backend releases the slot
750                // before deadpool drops this and creates a replacement.
751                let _ = socket2::SockRef::from(conn.underlying_tcp_stream())
752                    .shutdown(std::net::Shutdown::Both);
753                Err(e)
754            }
755        }
756    }
757
758    fn detach(&self, _conn: &mut ConnectionStream) {}
759}
760
761#[cfg(test)]
762mod tests {
763    use super::*;
764    use tokio::io::AsyncReadExt;
765    use tokio::net::TcpListener;
766
767    #[test]
768    fn test_socket_buffer_size_u32_rejects_oversized_values() {
769        let result = TcpManager::socket_buffer_size_u32(u32::MAX as usize + 1, "receive");
770
771        assert!(matches!(
772            result,
773            Err(ConnectionError::IoError(ref error))
774                if error.kind() == io::ErrorKind::InvalidInput
775                    && error.to_string().contains("receive socket buffer size")
776        ));
777    }
778
779    #[test]
780    fn test_tcp_manager_new_plain() {
781        let manager = TcpManager::new(
782            "news.example.com".to_string(),
783            119,
784            "TestServer".to_string(),
785            TcpManagerOptions {
786                username: Some("user".to_string()),
787                password: Some("pass".to_string()),
788                ..TcpManagerOptions::default()
789            },
790        )
791        .unwrap();
792
793        assert_eq!(manager.host, "news.example.com");
794        assert_eq!(manager.port, 119);
795        assert_eq!(manager.name, "TestServer");
796        assert_eq!(manager.username, Some("user".to_string()));
797        assert_eq!(manager.password, Some("pass".to_string()));
798        assert!(!manager.tls_config.use_tls);
799        assert!(manager.tls_manager.is_none());
800    }
801
802    #[test]
803    fn test_tcp_manager_new_without_auth() {
804        let manager = TcpManager::new(
805            "news.example.com".to_string(),
806            563,
807            "SecureServer".to_string(),
808            TcpManagerOptions::default(),
809        )
810        .unwrap();
811
812        assert_eq!(manager.host, "news.example.com");
813        assert_eq!(manager.port, 563);
814        assert_eq!(manager.name, "SecureServer");
815        assert!(manager.username.is_none());
816        assert!(manager.password.is_none());
817    }
818
819    #[test]
820    fn ip_literal_socket_addr_parses_ipv4_without_dns() {
821        let manager = TcpManager::new(
822            "127.0.0.1".to_string(),
823            119,
824            "IpBackend".to_string(),
825            TcpManagerOptions::default(),
826        )
827        .unwrap();
828
829        assert_eq!(
830            manager.ip_literal_socket_addr(),
831            Some(SocketAddr::from(([127, 0, 0, 1], 119)))
832        );
833    }
834
835    #[test]
836    fn ip_literal_socket_addr_parses_ipv6_without_dns() {
837        let manager = TcpManager::new(
838            "::1".to_string(),
839            563,
840            "IpBackend".to_string(),
841            TcpManagerOptions::default(),
842        )
843        .unwrap();
844
845        assert_eq!(
846            manager.ip_literal_socket_addr(),
847            Some(SocketAddr::from(([0, 0, 0, 0, 0, 0, 0, 1], 563)))
848        );
849    }
850
851    #[test]
852    fn ip_literal_socket_addr_leaves_hostnames_for_dns() {
853        let manager = TcpManager::new(
854            "news.example.com".to_string(),
855            119,
856            "DnsBackend".to_string(),
857            TcpManagerOptions::default(),
858        )
859        .unwrap();
860
861        assert_eq!(manager.ip_literal_socket_addr(), None);
862    }
863
864    #[test]
865    fn ipv6_network_unreachable_matches_error_kind_only_for_ipv6() {
866        let ipv6_addr = SocketAddr::from(([0, 0, 0, 0, 0, 0, 0, 1], 563));
867        let ipv4_addr = SocketAddr::from(([127, 0, 0, 1], 563));
868        let error = ConnectionError::IoError(io::Error::new(
869            io::ErrorKind::NetworkUnreachable,
870            "network unreachable",
871        ));
872
873        assert!(TcpManager::is_ipv6_network_unreachable(ipv6_addr, &error));
874        assert!(!TcpManager::is_ipv6_network_unreachable(ipv4_addr, &error));
875    }
876
877    #[tokio::test]
878    async fn remove_cached_ipv6_socket_addrs_keeps_only_ipv4_addresses() {
879        let manager = TcpManager::new(
880            "test.example.com".to_string(),
881            563,
882            "DnsBackend".to_string(),
883            TcpManagerOptions::default(),
884        )
885        .unwrap();
886        let ipv6_addr = SocketAddr::from(([0, 0, 0, 0, 0, 0, 0, 1], 563));
887        let ipv4_addr = SocketAddr::from(([127, 0, 0, 1], 563));
888        *manager.resolved_socket_addrs.write().await = Some(Arc::from([ipv6_addr, ipv4_addr]));
889
890        manager.remove_cached_ipv6_socket_addrs().await;
891
892        let cached_addrs = manager
893            .resolved_socket_addrs
894            .read()
895            .await
896            .as_ref()
897            .expect("cached addresses should remain initialized")
898            .clone();
899        assert_eq!(&*cached_addrs, &[ipv4_addr]);
900    }
901
902    #[tokio::test]
903    async fn remove_cached_ipv6_socket_addrs_clears_ipv6_only_cache() {
904        let manager = TcpManager::new(
905            "test.example.com".to_string(),
906            563,
907            "DnsBackend".to_string(),
908            TcpManagerOptions::default(),
909        )
910        .unwrap();
911        let ipv6_addr = SocketAddr::from(([0, 0, 0, 0, 0, 0, 0, 1], 563));
912        *manager.resolved_socket_addrs.write().await = Some(Arc::from([ipv6_addr]));
913
914        manager.remove_cached_ipv6_socket_addrs().await;
915
916        assert!(
917            manager.resolved_socket_addrs.read().await.is_none(),
918            "IPv6-only cached address list should be cleared instead of retained as an empty cache"
919        );
920    }
921
922    #[tokio::test]
923    async fn create_connected_tcp_stream_refreshes_dns_after_cached_addresses_fail() {
924        let live_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
925        let live_addr = live_listener.local_addr().unwrap();
926        let stale_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
927        let stale_addr = stale_listener.local_addr().unwrap();
928        drop(stale_listener);
929
930        let manager = TcpManager::new(
931            "localhost".to_string(),
932            live_addr.port(),
933            "DnsBackend".to_string(),
934            TcpManagerOptions::default(),
935        )
936        .unwrap();
937        *manager.resolved_socket_addrs.write().await = Some(Arc::from([stale_addr]));
938
939        let accept = tokio::spawn(async move { live_listener.accept().await.unwrap() });
940        let stream = manager.create_connected_tcp_stream().await.unwrap();
941        let _accepted = accept.await.unwrap();
942
943        assert_eq!(stream.peer_addr().unwrap(), live_addr);
944        let cached_addrs = manager
945            .resolved_socket_addrs
946            .read()
947            .await
948            .as_ref()
949            .expect("successful refresh should update cached addresses")
950            .clone();
951        assert!(cached_addrs.contains(&live_addr));
952    }
953
954    #[test]
955    fn test_tcp_manager_new_with_tls_disabled() {
956        let tls_config = TlsConfig::default(); // use_tls = false
957        let manager = TcpManager::new(
958            "news.example.com".to_string(),
959            119,
960            "PlainServer".to_string(),
961            TcpManagerOptions {
962                username: Some("user".to_string()),
963                password: Some("pass".to_string()),
964                tls_config: Some(tls_config),
965                ..TcpManagerOptions::default()
966            },
967        )
968        .unwrap();
969
970        assert_eq!(manager.host, "news.example.com");
971        assert_eq!(manager.port, 119);
972        assert!(!manager.tls_config.use_tls);
973        assert!(manager.tls_manager.is_none());
974    }
975
976    #[test]
977    fn test_tcp_manager_new_with_tls_enabled() {
978        let tls_config = TlsConfig {
979            use_tls: true,
980            tls_verify_cert: true,
981            tls_cert_path: None,
982        };
983        let manager = TcpManager::new(
984            "secure.example.com".to_string(),
985            563,
986            "SecureServer".to_string(),
987            TcpManagerOptions {
988                username: Some("user".to_string()),
989                password: Some("pass".to_string()),
990                tls_config: Some(tls_config),
991                ..TcpManagerOptions::default()
992            },
993        )
994        .unwrap();
995
996        assert_eq!(manager.host, "secure.example.com");
997        assert_eq!(manager.port, 563);
998        assert!(manager.tls_config.use_tls);
999        assert!(manager.tls_manager.is_some());
1000    }
1001
1002    #[test]
1003    fn test_tcp_manager_clone() {
1004        let manager = TcpManager::new(
1005            "news.example.com".to_string(),
1006            119,
1007            "TestServer".to_string(),
1008            TcpManagerOptions {
1009                username: Some("user".to_string()),
1010                password: Some("pass".to_string()),
1011                ..TcpManagerOptions::default()
1012            },
1013        )
1014        .unwrap();
1015
1016        let cloned = manager.clone();
1017        assert_eq!(cloned.host, manager.host);
1018        assert_eq!(cloned.port, manager.port);
1019        assert_eq!(cloned.name, manager.name);
1020        assert_eq!(cloned.username, manager.username);
1021        assert_eq!(cloned.password, manager.password);
1022        assert!(Arc::ptr_eq(
1023            &cloned.resolved_socket_addrs,
1024            &manager.resolved_socket_addrs
1025        ));
1026        assert!(Arc::ptr_eq(
1027            &cloned.next_resolved_socket_addr,
1028            &manager.next_resolved_socket_addr
1029        ));
1030    }
1031
1032    #[test]
1033    fn test_tcp_manager_debug_format() {
1034        let manager = TcpManager::new(
1035            "news.example.com".to_string(),
1036            119,
1037            "TestServer".to_string(),
1038            TcpManagerOptions {
1039                username: Some("user".to_string()),
1040                password: Some("pass".to_string()),
1041                ..TcpManagerOptions::default()
1042            },
1043        )
1044        .unwrap();
1045
1046        let debug_str = format!("{manager:?}");
1047        assert!(debug_str.contains("TcpManager"));
1048        assert!(debug_str.contains("news.example.com"));
1049        assert!(debug_str.contains("119"));
1050    }
1051
1052    #[tokio::test]
1053    async fn create_optimized_stream_tries_next_resolved_address_after_connect_error() {
1054        let unavailable_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1055        let unavailable_addr = unavailable_listener.local_addr().unwrap();
1056        drop(unavailable_listener);
1057
1058        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1059        let available_addr = listener.local_addr().unwrap();
1060        let accept_task = tokio::spawn(async move {
1061            let (_stream, _) = listener.accept().await.unwrap();
1062        });
1063
1064        let manager = TcpManager::new(
1065            "test.example.com".to_string(),
1066            available_addr.port(),
1067            "FallbackBackend".to_string(),
1068            TcpManagerOptions::default(),
1069        )
1070        .unwrap();
1071        *manager.resolved_socket_addrs.write().await =
1072            Some(Arc::from([unavailable_addr, available_addr]));
1073
1074        let stream = manager
1075            .create_optimized_stream()
1076            .await
1077            .expect("second resolved address should be tried after first connect error");
1078
1079        assert_eq!(stream.connection_type(), "TCP");
1080        accept_task.await.unwrap();
1081    }
1082
1083    #[test]
1084    fn test_tcp_manager_with_tls_manager_is_some() {
1085        let tls_config = TlsConfig {
1086            use_tls: true,
1087            tls_verify_cert: false,
1088            tls_cert_path: None,
1089        };
1090        let manager = TcpManager::new(
1091            "secure.example.com".to_string(),
1092            563,
1093            "SecureServer".to_string(),
1094            TcpManagerOptions {
1095                tls_config: Some(tls_config),
1096                ..TcpManagerOptions::default()
1097            },
1098        )
1099        .unwrap();
1100
1101        assert!(manager.tls_manager.is_some());
1102
1103        // Verify TLS manager is an Arc (cheap clone)
1104        let arc_clone = manager.tls_manager.as_ref().unwrap().clone();
1105        assert!(Arc::ptr_eq(
1106            manager.tls_manager.as_ref().unwrap(),
1107            &arc_clone
1108        ));
1109    }
1110
1111    #[test]
1112    fn test_tcp_manager_with_tls_cert_path() {
1113        let tls_config = TlsConfig {
1114            use_tls: true,
1115            tls_verify_cert: true,
1116            tls_cert_path: Some("/path/to/ca.pem".to_string()),
1117        };
1118
1119        // This will fail due to missing file, but tests the construction path
1120        let result = TcpManager::new(
1121            "secure.example.com".to_string(),
1122            563,
1123            "SecureServer".to_string(),
1124            TcpManagerOptions {
1125                tls_config: Some(tls_config),
1126                ..TcpManagerOptions::default()
1127            },
1128        );
1129
1130        // Should fail because cert file doesn't exist
1131        assert!(result.is_err());
1132    }
1133
1134    #[tokio::test]
1135    async fn mode_reader_negotiation_reads_split_setup_reply() {
1136        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1137        let addr = listener.local_addr().unwrap();
1138
1139        tokio::spawn(async move {
1140            let (mut stream, _) = listener.accept().await.unwrap();
1141            let mut command = [0u8; 13];
1142            stream.read_exact(&mut command).await.unwrap();
1143            assert_eq!(&command, b"MODE READER\r\n");
1144            stream.write_all(b"20").await.unwrap();
1145            stream.write_all(b"0 Posting allowed\r\n").await.unwrap();
1146        });
1147
1148        let manager = TcpManager::new(
1149            addr.ip().to_string(),
1150            addr.port(),
1151            "SplitSetup".to_string(),
1152            TcpManagerOptions::default(),
1153        )
1154        .unwrap();
1155        let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
1156        let mut stream = ConnectionStream::plain(tcp_stream);
1157        let mut buffer = [0u8; 64];
1158
1159        manager
1160            .negotiate_mode_reader(&mut stream, &mut buffer)
1161            .await
1162            .expect("split MODE READER setup reply should be accepted");
1163    }
1164
1165    #[tokio::test]
1166    async fn compression_negotiation_reads_split_unsupported_reply() {
1167        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1168        let addr = listener.local_addr().unwrap();
1169
1170        tokio::spawn(async move {
1171            let (mut stream, _) = listener.accept().await.unwrap();
1172            let mut command = [0u8; 18];
1173            stream.read_exact(&mut command).await.unwrap();
1174            assert_eq!(&command, crate::protocol::COMPRESS_DEFLATE);
1175            stream.write_all(b"50").await.unwrap();
1176            stream.write_all(b"0 Not supported\r\n").await.unwrap();
1177        });
1178
1179        let manager = TcpManager::new(
1180            addr.ip().to_string(),
1181            addr.port(),
1182            "SplitCompression".to_string(),
1183            TcpManagerOptions {
1184                compress: None,
1185                ..TcpManagerOptions::default()
1186            },
1187        )
1188        .unwrap();
1189        let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
1190        let mut stream = ConnectionStream::plain(tcp_stream);
1191        let mut buffer = [0u8; 64];
1192
1193        let enabled = manager
1194            .negotiate_compression(&mut stream, &mut buffer)
1195            .await
1196            .expect("split unsupported compression reply should be accepted");
1197
1198        assert!(!enabled);
1199    }
1200
1201    #[tokio::test]
1202    async fn auto_compression_serializes_and_remembers_unsupported_backend() {
1203        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1204        let addr = listener.local_addr().unwrap();
1205        let compress_commands = Arc::new(AtomicUsize::new(0));
1206        let server_commands = compress_commands.clone();
1207
1208        tokio::spawn(async move {
1209            for _ in 0..2 {
1210                let (mut stream, _) = listener.accept().await.unwrap();
1211                let server_commands = server_commands.clone();
1212                tokio::spawn(async move {
1213                    let mut command = [0u8; 18];
1214                    let read = tokio::time::timeout(
1215                        std::time::Duration::from_millis(100),
1216                        stream.read_exact(&mut command),
1217                    )
1218                    .await;
1219                    if read.is_err() {
1220                        return;
1221                    }
1222                    read.unwrap().unwrap();
1223                    if command == crate::protocol::COMPRESS_DEFLATE {
1224                        server_commands.fetch_add(1, Ordering::SeqCst);
1225                        stream.write_all(b"500 Not supported\r\n").await.unwrap();
1226                    }
1227                });
1228            }
1229        });
1230
1231        let manager = TcpManager::new(
1232            addr.ip().to_string(),
1233            addr.port(),
1234            "CachedUnsupportedCompression".to_string(),
1235            TcpManagerOptions {
1236                compress: None,
1237                ..TcpManagerOptions::default()
1238            },
1239        )
1240        .unwrap();
1241
1242        let mut tasks = Vec::new();
1243        for _ in 0..2 {
1244            let manager = manager.clone();
1245            tasks.push(tokio::spawn(async move {
1246                let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
1247                let mut stream = ConnectionStream::plain(tcp_stream);
1248                let mut buffer = [0u8; 64];
1249
1250                let enabled = manager
1251                    .negotiate_compression(&mut stream, &mut buffer)
1252                    .await
1253                    .expect("unsupported compression should fall back in auto mode");
1254
1255                assert!(!enabled);
1256            }));
1257        }
1258
1259        for task in tasks {
1260            task.await.unwrap();
1261        }
1262
1263        assert_eq!(
1264            compress_commands.load(Ordering::SeqCst),
1265            1,
1266            "auto mode should remember unsupported COMPRESS DEFLATE"
1267        );
1268    }
1269}