1use 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
23pub 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#[derive(Debug, Clone)]
46pub struct TcpManagerOptions {
47 pub username: Option<String>,
48 pub password: Option<String>,
49 pub tls_config: Option<TlsConfig>,
50 pub recv_buffer_size: usize,
52 pub send_buffer_size: usize,
54 pub compress: Option<bool>,
56 pub compress_level: Option<u32>,
58 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#[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 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 pub(crate) compress: Option<bool>,
97 pub(crate) compress_level: Option<u32>,
99 compression_support: Arc<Mutex<CompressionSupportState>>,
100 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 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 let socket = if socket_addr.is_ipv4() {
265 tokio::net::TcpSocket::new_v4()?
266 } else {
267 tokio::net::TcpSocket::new_v6()?
268 };
269
270 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 let tcp_stream = socket.connect(socket_addr).await?;
287
288 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 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 if self.tls_config.use_tls {
377 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
399impl TcpManager {
404 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 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 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 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 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 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 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 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 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 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
713impl 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 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(); 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 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 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 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}