Skip to main content

nntp_proxy/
network.rs

1//! Network socket optimization utilities
2//!
3//! This module provides utilities for optimizing TCP socket performance
4//! for high-throughput NNTP transfers.
5//!
6//! Includes trait-based optimization strategies for different connection types
7//! (plain TCP, TLS) with configurable buffer sizes.
8
9use crate::stream::ConnectionStream;
10use anyhow::{Context, Result};
11use socket2::SockRef;
12use std::time::Duration;
13use tokio::net::TcpStream;
14use tracing::debug;
15
16/// `SO_LINGER` timeout - prevents indefinite blocking on socket close
17const LINGER_TIMEOUT: Duration = Duration::from_secs(5);
18
19/// `TCP_USER_TIMEOUT` - faster dead connection detection on Linux
20#[cfg(target_os = "linux")]
21const TCP_USER_TIMEOUT: Duration = Duration::from_secs(30);
22
23/// `IP_TOS` value for throughput optimization
24#[cfg(target_os = "linux")]
25const TOS_THROUGHPUT: u32 = 0x08;
26
27/// Trait for network optimization strategies
28pub trait NetworkOptimizer {
29    /// Apply optimizations to improve network performance
30    ///
31    /// # Errors
32    /// Returns any socket option error reported by the operating system while
33    /// applying the optimization strategy.
34    fn optimize(&self) -> Result<()>;
35
36    /// Get a description of the optimization strategy
37    fn description(&self) -> &'static str;
38}
39
40/// Apply core TCP optimizations to a socket reference
41fn apply_core_optimizations(
42    sock_ref: &SockRef,
43    recv_buffer_size: usize,
44    send_buffer_size: usize,
45) -> Result<()> {
46    if recv_buffer_size > 0 {
47        sock_ref
48            .set_recv_buffer_size(recv_buffer_size)
49            .context("Failed to set TCP receive buffer size")?;
50    }
51
52    if send_buffer_size > 0 {
53        sock_ref
54            .set_send_buffer_size(send_buffer_size)
55            .context("Failed to set TCP send buffer size")?;
56    }
57
58    sock_ref
59        .set_linger(Some(LINGER_TIMEOUT))
60        .context("Failed to set SO_LINGER timeout")?;
61
62    sock_ref
63        .set_tcp_nodelay(true)
64        .context("Failed to set TCP_NODELAY")?;
65
66    Ok(())
67}
68
69/// Apply Linux-specific TCP optimizations (best-effort)
70#[cfg(target_os = "linux")]
71fn apply_linux_optimizations(sock_ref: &SockRef, context: &str) {
72    [
73        (
74            "TCP_USER_TIMEOUT",
75            sock_ref.set_tcp_user_timeout(Some(TCP_USER_TIMEOUT)),
76        ),
77        ("IP_TOS", sock_ref.set_tos_v4(TOS_THROUGHPUT)),
78    ]
79    .into_iter()
80    .filter_map(|(name, result)| result.err().map(|e| (name, e)))
81    .for_each(|(name, err)| {
82        debug!("Failed to set {} on {}: {}", name, context, err);
83    });
84}
85
86/// Get platform-specific optimization description
87const fn platform_optimization_desc() -> &'static str {
88    match () {
89        #[cfg(target_os = "linux")]
90        () => ", tcp_user_timeout=30s, tos=0x08",
91        #[cfg(target_os = "windows")]
92        () => " (Windows)",
93        #[cfg(not(any(target_os = "linux", target_os = "windows")))]
94        () => "",
95    }
96}
97
98/// TCP-specific optimizations for high-throughput scenarios
99pub struct TcpOptimizer<'a> {
100    stream: &'a TcpStream,
101    recv_buffer_size: usize,
102    send_buffer_size: usize,
103}
104
105impl<'a> TcpOptimizer<'a> {
106    /// Create a new TCP optimizer with default high-throughput settings
107    pub const fn new(stream: &'a TcpStream) -> Self {
108        Self {
109            stream,
110            recv_buffer_size: crate::constants::socket::HIGH_THROUGHPUT_RECV_BUFFER,
111            send_buffer_size: crate::constants::socket::HIGH_THROUGHPUT_SEND_BUFFER,
112        }
113    }
114
115    /// Create optimizer with custom buffer sizes using builder pattern
116    pub const fn with_buffer_sizes(
117        stream: &'a TcpStream,
118        recv_size: usize,
119        send_size: usize,
120    ) -> Self {
121        Self {
122            stream,
123            recv_buffer_size: recv_size,
124            send_buffer_size: send_size,
125        }
126    }
127}
128
129impl NetworkOptimizer for TcpOptimizer<'_> {
130    fn optimize(&self) -> Result<()> {
131        let sock_ref = SockRef::from(self.stream);
132
133        // Core optimizations (required)
134        apply_core_optimizations(&sock_ref, self.recv_buffer_size, self.send_buffer_size)
135            .context("Failed to apply core TCP optimizations")?;
136
137        // Platform-specific optimizations (best-effort)
138        #[cfg(target_os = "linux")]
139        apply_linux_optimizations(&sock_ref, "TCP stream");
140
141        debug!(
142            "Applied TCP optimizations: recv_buffer={}, send_buffer={}, linger={}s, nodelay=true{}",
143            self.recv_buffer_size,
144            self.send_buffer_size,
145            LINGER_TIMEOUT.as_secs(),
146            platform_optimization_desc()
147        );
148
149        Ok(())
150    }
151
152    fn description(&self) -> &'static str {
153        "TCP high-throughput optimization"
154    }
155}
156
157/// High-level optimizer that works with `ConnectionStream`
158pub struct ConnectionOptimizer<'a> {
159    stream: &'a ConnectionStream,
160    recv_buffer_size: Option<usize>,
161    send_buffer_size: Option<usize>,
162}
163
164impl<'a> ConnectionOptimizer<'a> {
165    /// Create a new connection optimizer with default buffer sizes
166    pub const fn new(stream: &'a ConnectionStream) -> Self {
167        Self {
168            stream,
169            recv_buffer_size: None,
170            send_buffer_size: None,
171        }
172    }
173
174    /// Create a connection optimizer with custom buffer sizes
175    pub const fn with_buffer_sizes(
176        stream: &'a ConnectionStream,
177        recv_size: usize,
178        send_size: usize,
179    ) -> Self {
180        Self {
181            stream,
182            recv_buffer_size: Some(recv_size),
183            send_buffer_size: Some(send_size),
184        }
185    }
186}
187
188impl NetworkOptimizer for ConnectionOptimizer<'_> {
189    fn optimize(&self) -> Result<()> {
190        // Use functional pattern matching to create and optimize in one step
191        let optimize_fn = |desc: &str, result: Result<()>| {
192            debug!("Using {}", desc);
193            result
194        };
195
196        // Get the underlying TCP stream for optimization regardless of compression layer
197        let tcp = self.stream.underlying_tcp_stream();
198        if let (Some(recv), Some(send)) = (self.recv_buffer_size, self.send_buffer_size) {
199            let desc = if self.stream.is_encrypted() {
200                "TLS optimization via underlying TCP stream with custom buffers"
201            } else {
202                "TCP high-throughput optimization with custom buffers"
203            };
204            optimize_fn(
205                desc,
206                TcpOptimizer::with_buffer_sizes(tcp, recv, send).optimize(),
207            )
208        } else {
209            let desc = if self.stream.is_encrypted() {
210                "TLS optimization via underlying TCP stream"
211            } else {
212                "TCP high-throughput optimization"
213            };
214            optimize_fn(desc, TcpOptimizer::new(tcp).optimize())
215        }
216    }
217
218    fn description(&self) -> &'static str {
219        if self.stream.is_encrypted() {
220            "Connection-level TLS optimization"
221        } else {
222            "Connection-level TCP optimization"
223        }
224    }
225}
226
227#[cfg(test)]
228mod tests {
229    use super::*;
230    use crate::constants::socket::{HIGH_THROUGHPUT_RECV_BUFFER, HIGH_THROUGHPUT_SEND_BUFFER};
231    use tokio::net::TcpListener;
232
233    #[test]
234    fn test_constants() {
235        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER, 16 * 1024 * 1024);
236        assert_eq!(HIGH_THROUGHPUT_SEND_BUFFER, 16 * 1024 * 1024);
237    }
238
239    #[test]
240    fn test_buffer_size_is_reasonable() {
241        // Buffer sizes should be large but not excessive
242        const _: () = assert!(HIGH_THROUGHPUT_RECV_BUFFER >= 1024 * 1024); // At least 1MB
243        const _: () = assert!(HIGH_THROUGHPUT_RECV_BUFFER <= 128 * 1024 * 1024); // At most 128MB
244        const _: () = assert!(HIGH_THROUGHPUT_SEND_BUFFER >= 1024 * 1024);
245        const _: () = assert!(HIGH_THROUGHPUT_SEND_BUFFER <= 128 * 1024 * 1024);
246    }
247
248    #[test]
249    fn test_buffer_sizes_are_equal() {
250        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER, HIGH_THROUGHPUT_SEND_BUFFER);
251    }
252
253    #[test]
254    fn test_buffer_sizes_are_power_of_two_or_multiple() {
255        let size = HIGH_THROUGHPUT_RECV_BUFFER;
256        assert_eq!(size % (1024 * 1024), 0);
257    }
258
259    #[tokio::test]
260    async fn test_connection_optimizer() {
261        let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
262        let addr = listener.local_addr().unwrap();
263
264        let tcp_stream = std::net::TcpStream::connect(addr).unwrap();
265        tcp_stream.set_nonblocking(true).unwrap();
266        let tokio_stream = TcpStream::from_std(tcp_stream).unwrap();
267
268        let conn_stream = ConnectionStream::plain(tokio_stream);
269
270        let optimizer = ConnectionOptimizer::new(&conn_stream);
271        let result = optimizer.optimize();
272        assert!(result.is_ok());
273    }
274
275    #[test]
276    fn test_buffer_size_calculation() {
277        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER, 16 * 1024 * 1024);
278        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER, 16_777_216);
279        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER / 1024, 16384); // KB
280        assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER / (1024 * 1024), 16); // MB
281    }
282
283    #[test]
284    fn test_buffer_size_for_large_articles() {
285        let typical_large_article = 10 * 1024 * 1024; // 10MB
286        let very_large_article = 100 * 1024 * 1024; // 100MB
287        assert!(HIGH_THROUGHPUT_RECV_BUFFER > typical_large_article);
288        assert!(HIGH_THROUGHPUT_RECV_BUFFER < very_large_article);
289    }
290
291    #[tokio::test]
292    async fn test_tcp_optimizer() {
293        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
294        let addr = listener.local_addr().unwrap();
295
296        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
297        let optimizer = TcpOptimizer::new(&stream);
298
299        assert_eq!(optimizer.description(), "TCP high-throughput optimization");
300
301        // Should not panic - actual socket optimization might fail in test environment
302        let _ = optimizer.optimize();
303    }
304
305    #[tokio::test]
306    async fn test_connection_optimizer_tcp() {
307        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
308        let addr = listener.local_addr().unwrap();
309
310        let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
311        let connection_stream = ConnectionStream::plain(tcp_stream);
312        let optimizer = ConnectionOptimizer::new(&connection_stream);
313
314        assert_eq!(optimizer.description(), "Connection-level TCP optimization");
315
316        // Should not panic - actual socket optimization might fail in test environment
317        let _ = optimizer.optimize();
318    }
319
320    #[tokio::test]
321    async fn test_connection_optimizer_trait_usage() {
322        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
323        let addr = listener.local_addr().unwrap();
324
325        let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
326        let connection_stream = ConnectionStream::plain(tcp_stream);
327
328        // Test that ConnectionOptimizer implements NetworkOptimizer trait
329        let optimizer: Box<dyn NetworkOptimizer> =
330            Box::new(ConnectionOptimizer::new(&connection_stream));
331
332        assert_eq!(optimizer.description(), "Connection-level TCP optimization");
333        let _ = optimizer.optimize();
334    }
335
336    #[tokio::test]
337    async fn test_connection_optimizer_with_custom_buffers() {
338        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
339        let addr = listener.local_addr().unwrap();
340
341        let tcp_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
342        let connection_stream = ConnectionStream::plain(tcp_stream);
343        let optimizer = ConnectionOptimizer::with_buffer_sizes(&connection_stream, 4096, 8192);
344
345        assert_eq!(optimizer.description(), "Connection-level TCP optimization");
346        let _ = optimizer.optimize();
347    }
348
349    #[tokio::test]
350    async fn test_optimizer_creation() {
351        // Test that we can create optimizers with tokio streams
352        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
353        let addr = listener.local_addr().unwrap();
354
355        let tokio_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
356
357        let optimizer = TcpOptimizer::new(&tokio_stream);
358        assert_eq!(
359            optimizer.recv_buffer_size,
360            crate::constants::socket::HIGH_THROUGHPUT_RECV_BUFFER
361        );
362        assert_eq!(
363            optimizer.send_buffer_size,
364            crate::constants::socket::HIGH_THROUGHPUT_SEND_BUFFER
365        );
366    }
367
368    #[tokio::test]
369    async fn test_custom_buffer_sizes() {
370        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
371        let addr = listener.local_addr().unwrap();
372
373        let tokio_stream = tokio::net::TcpStream::connect(addr).await.unwrap();
374
375        let optimizer = TcpOptimizer::with_buffer_sizes(&tokio_stream, 1024, 2048);
376        assert_eq!(optimizer.recv_buffer_size, 1024);
377        assert_eq!(optimizer.send_buffer_size, 2048);
378    }
379
380    // Platform-specific description tests
381
382    #[test]
383    fn test_platform_optimization_desc() {
384        let desc = platform_optimization_desc();
385        // Check that it returns a valid string (platform-specific content)
386        #[cfg(target_os = "linux")]
387        assert!(desc.contains("tcp_user_timeout"));
388        #[cfg(target_os = "windows")]
389        assert!(desc.contains("Windows"));
390        // Just verify it doesn't panic on other platforms
391        #[cfg(not(any(target_os = "linux", target_os = "windows")))]
392        let _ = desc;
393    }
394
395    // Constructor tests
396
397    #[tokio::test]
398    async fn test_tcp_optimizer_new_uses_defaults() {
399        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
400        let addr = listener.local_addr().unwrap();
401        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
402
403        let optimizer = TcpOptimizer::new(&stream);
404
405        assert_eq!(
406            optimizer.recv_buffer_size,
407            crate::constants::socket::HIGH_THROUGHPUT_RECV_BUFFER
408        );
409        assert_eq!(
410            optimizer.send_buffer_size,
411            crate::constants::socket::HIGH_THROUGHPUT_SEND_BUFFER
412        );
413    }
414
415    #[tokio::test]
416    async fn test_tcp_optimizer_with_custom_buffers() {
417        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
418        let addr = listener.local_addr().unwrap();
419        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
420
421        let recv_size = 512 * 1024;
422        let send_size = 1024 * 1024;
423        let optimizer = TcpOptimizer::with_buffer_sizes(&stream, recv_size, send_size);
424
425        assert_eq!(optimizer.recv_buffer_size, recv_size);
426        assert_eq!(optimizer.send_buffer_size, send_size);
427    }
428
429    // Description tests
430
431    #[tokio::test]
432    async fn test_tcp_optimizer_description() {
433        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
434        let addr = listener.local_addr().unwrap();
435        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
436
437        let optimizer = TcpOptimizer::new(&stream);
438        assert_eq!(optimizer.description(), "TCP high-throughput optimization");
439    }
440
441    #[tokio::test]
442    async fn test_connection_optimizer_description_tcp() {
443        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
444        let addr = listener.local_addr().unwrap();
445        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
446        let stream = ConnectionStream::plain(tcp);
447
448        let optimizer = ConnectionOptimizer::new(&stream);
449        assert_eq!(optimizer.description(), "Connection-level TCP optimization");
450    }
451
452    // ConnectionOptimizer buffer configuration tests
453
454    #[tokio::test]
455    async fn test_connection_optimizer_new_no_custom_buffers() {
456        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
457        let addr = listener.local_addr().unwrap();
458        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
459        let stream = ConnectionStream::plain(tcp);
460
461        let optimizer = ConnectionOptimizer::new(&stream);
462        assert!(optimizer.recv_buffer_size.is_none());
463        assert!(optimizer.send_buffer_size.is_none());
464    }
465
466    #[tokio::test]
467    async fn test_connection_optimizer_with_custom_buffers_sets_sizes() {
468        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
469        let addr = listener.local_addr().unwrap();
470        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
471        let stream = ConnectionStream::plain(tcp);
472
473        let recv = 16384;
474        let send = 32768;
475        let optimizer = ConnectionOptimizer::with_buffer_sizes(&stream, recv, send);
476        assert_eq!(optimizer.recv_buffer_size, Some(recv));
477        assert_eq!(optimizer.send_buffer_size, Some(send));
478    }
479
480    // Edge case tests
481
482    #[tokio::test]
483    async fn test_tcp_optimizer_with_zero_buffers() {
484        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
485        let addr = listener.local_addr().unwrap();
486        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
487
488        // Zero buffer sizes leave the OS defaults in place.
489        let optimizer = TcpOptimizer::with_buffer_sizes(&stream, 0, 0);
490        assert_eq!(optimizer.recv_buffer_size, 0);
491        assert_eq!(optimizer.send_buffer_size, 0);
492        optimizer.optimize().unwrap();
493    }
494
495    #[tokio::test]
496    async fn test_tcp_optimizer_with_very_large_buffers() {
497        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
498        let addr = listener.local_addr().unwrap();
499        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
500
501        // Very large buffer sizes (system may reject these, but we test the constructor)
502        let large_size = 128 * 1024 * 1024; // 128MB
503        let optimizer = TcpOptimizer::with_buffer_sizes(&stream, large_size, large_size);
504        assert_eq!(optimizer.recv_buffer_size, large_size);
505        assert_eq!(optimizer.send_buffer_size, large_size);
506    }
507
508    #[tokio::test]
509    async fn test_tcp_optimizer_with_asymmetric_buffers() {
510        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
511        let addr = listener.local_addr().unwrap();
512        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
513
514        // Asymmetric buffers (common for download-heavy workloads)
515        let recv_size = 32 * 1024 * 1024; // 32MB
516        let send_size = 1024 * 1024; // 1MB
517        let optimizer = TcpOptimizer::with_buffer_sizes(&stream, recv_size, send_size);
518        assert_eq!(optimizer.recv_buffer_size, recv_size);
519        assert_eq!(optimizer.send_buffer_size, send_size);
520    }
521
522    // NetworkOptimizer trait tests
523
524    #[tokio::test]
525    async fn test_tcp_optimizer_implements_network_optimizer() {
526        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
527        let addr = listener.local_addr().unwrap();
528        let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
529
530        let optimizer = TcpOptimizer::new(&stream);
531        // Can be used as a trait object
532        let _: &dyn NetworkOptimizer = &optimizer;
533    }
534
535    #[tokio::test]
536    async fn test_connection_optimizer_implements_network_optimizer() {
537        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
538        let addr = listener.local_addr().unwrap();
539        let tcp = tokio::net::TcpStream::connect(addr).await.unwrap();
540        let stream = ConnectionStream::plain(tcp);
541
542        let optimizer = ConnectionOptimizer::new(&stream);
543        // Can be used as a trait object
544        let _: &dyn NetworkOptimizer = &optimizer;
545    }
546
547    // Constants tests
548
549    #[test]
550    fn test_linger_timeout_constant() {
551        assert_eq!(LINGER_TIMEOUT, Duration::from_secs(5));
552    }
553
554    #[test]
555    #[cfg(target_os = "linux")]
556    fn test_tcp_user_timeout_constant() {
557        assert_eq!(TCP_USER_TIMEOUT, Duration::from_secs(30));
558    }
559
560    #[test]
561    #[cfg(target_os = "linux")]
562    fn test_tos_throughput_constant() {
563        assert_eq!(TOS_THROUGHPUT, 0x08);
564    }
565}