1use crate::stream::ConnectionStream;
10use anyhow::{Context, Result};
11use socket2::SockRef;
12use std::time::Duration;
13use tokio::net::TcpStream;
14use tracing::debug;
15
16const LINGER_TIMEOUT: Duration = Duration::from_secs(5);
18
19#[cfg(target_os = "linux")]
21const TCP_USER_TIMEOUT: Duration = Duration::from_secs(30);
22
23#[cfg(target_os = "linux")]
25const TOS_THROUGHPUT: u32 = 0x08;
26
27pub trait NetworkOptimizer {
29 fn optimize(&self) -> Result<()>;
35
36 fn description(&self) -> &'static str;
38}
39
40fn 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#[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
86const 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
98pub struct TcpOptimizer<'a> {
100 stream: &'a TcpStream,
101 recv_buffer_size: usize,
102 send_buffer_size: usize,
103}
104
105impl<'a> TcpOptimizer<'a> {
106 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 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 apply_core_optimizations(&sock_ref, self.recv_buffer_size, self.send_buffer_size)
135 .context("Failed to apply core TCP optimizations")?;
136
137 #[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
157pub 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 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 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 let optimize_fn = |desc: &str, result: Result<()>| {
192 debug!("Using {}", desc);
193 result
194 };
195
196 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 const _: () = assert!(HIGH_THROUGHPUT_RECV_BUFFER >= 1024 * 1024); const _: () = assert!(HIGH_THROUGHPUT_RECV_BUFFER <= 128 * 1024 * 1024); 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); assert_eq!(HIGH_THROUGHPUT_RECV_BUFFER / (1024 * 1024), 16); }
282
283 #[test]
284 fn test_buffer_size_for_large_articles() {
285 let typical_large_article = 10 * 1024 * 1024; let very_large_article = 100 * 1024 * 1024; 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 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 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 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 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 #[test]
383 fn test_platform_optimization_desc() {
384 let desc = platform_optimization_desc();
385 #[cfg(target_os = "linux")]
387 assert!(desc.contains("tcp_user_timeout"));
388 #[cfg(target_os = "windows")]
389 assert!(desc.contains("Windows"));
390 #[cfg(not(any(target_os = "linux", target_os = "windows")))]
392 let _ = desc;
393 }
394
395 #[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 #[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 #[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 #[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 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 let large_size = 128 * 1024 * 1024; 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 let recv_size = 32 * 1024 * 1024; let send_size = 1024 * 1024; 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 #[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 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 let _: &dyn NetworkOptimizer = &optimizer;
545 }
546
547 #[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}