1use std::net::{IpAddr, Ipv6Addr, SocketAddr};
2
3use tokio::net::TcpStream;
4
5use crate::{BoxStream, ConnectError, TargetAddr, TargetHost};
6
7pub fn is_reserved_or_private_ip(ip: &IpAddr) -> bool {
28 match ip {
29 IpAddr::V4(v4) => {
30 v4.is_loopback()
31 || v4.is_link_local()
32 || v4.is_private()
33 || v4.is_unspecified()
34 || v4.is_multicast()
35 || v4.is_broadcast()
36 || is_v4_documentation(v4)
37 || is_v4_benchmarking(v4)
38 || is_v4_reserved(v4)
39 || is_v4_this_network(v4)
40 }
41 IpAddr::V6(v6) => {
42 if let Some(v4) = v6.to_ipv4_mapped() {
46 return is_reserved_or_private_ip(&IpAddr::V4(v4));
47 }
48 v6.is_loopback()
49 || v6.is_unspecified()
50 || v6.is_multicast()
51 || is_v6_documentation(v6)
52 || is_unicast_link_local_v6(v6)
53 || is_unique_local_v6(v6)
54 || is_v6_discard_prefix(v6)
55 }
56 }
57}
58
59fn is_unique_local_v6(ip: &Ipv6Addr) -> bool {
61 let octets = ip.octets();
62 (octets[0] & 0xfe) == 0xfc
63}
64
65fn is_unicast_link_local_v6(ip: &Ipv6Addr) -> bool {
67 let octets = ip.octets();
68 octets[0] == 0xfe && (octets[1] & 0xc0) == 0x80
69}
70
71fn is_v6_discard_prefix(ip: &Ipv6Addr) -> bool {
73 let octets = ip.octets();
74 octets[0] == 0x01 && octets[1..8].iter().all(|b| *b == 0)
75}
76
77fn is_v4_this_network(ip: &std::net::Ipv4Addr) -> bool {
79 ip.octets()[0] == 0
80}
81
82fn is_v4_documentation(ip: &std::net::Ipv4Addr) -> bool {
86 let octets = ip.octets();
87 matches!(
88 octets,
89 [192, 0, 2, _] | [198, 51, 100, _] | [203, 0, 113, _] | [192, 88, 99, _]
90 )
91}
92
93fn is_v4_benchmarking(ip: &std::net::Ipv4Addr) -> bool {
95 let octets = ip.octets();
96 octets[0] == 198 && (octets[1] == 18 || octets[1] == 19)
97}
98
99fn is_v4_reserved(ip: &std::net::Ipv4Addr) -> bool {
103 ip.octets()[0] >= 240
104}
105
106fn is_v6_documentation(ip: &Ipv6Addr) -> bool {
108 let octets = ip.octets();
109 octets[0] == 0x20 && octets[1] == 0x01 && octets[2] == 0x0d && octets[3] == 0xb8
110}
111
112pub fn is_dns_rebinding_risk(ip: &IpAddr) -> bool {
114 is_reserved_or_private_ip(ip)
115}
116
117#[trait_variant::make(Connector: Send)]
119pub trait LocalConnector {
120 async fn connect(&self, target: &TargetAddr) -> Result<BoxStream, ConnectError>;
121}
122
123pub struct DirectConnector;
125
126#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
132pub struct ConnectionMetadata {
133 local_addr: Option<SocketAddr>,
134 peer_addr: Option<SocketAddr>,
135}
136
137impl ConnectionMetadata {
138 pub fn local_addr(&self) -> Option<SocketAddr> {
139 self.local_addr
140 }
141
142 pub fn peer_addr(&self) -> Option<SocketAddr> {
143 self.peer_addr
144 }
145}
146
147#[derive(Debug, Clone)]
149pub struct ConnectOptions {
150 pub local_bind: Option<SocketAddr>,
151 pub enforce_dns_rebinding_check: bool,
154 pub enforce_literal_ip_check: bool,
157}
158
159impl Default for ConnectOptions {
160 fn default() -> Self {
161 Self {
162 local_bind: None,
163 enforce_dns_rebinding_check: true,
164 enforce_literal_ip_check: false,
165 }
166 }
167}
168
169impl DirectConnector {
170 pub async fn connect_with_options(
171 &self,
172 target: &TargetAddr,
173 options: &ConnectOptions,
174 ) -> Result<BoxStream, ConnectError> {
175 self.connect_with_options_and_metadata(target, options)
176 .await
177 .map(|(stream, _)| stream)
178 }
179
180 pub async fn connect_with_options_and_metadata(
182 &self,
183 target: &TargetAddr,
184 options: &ConnectOptions,
185 ) -> Result<(BoxStream, ConnectionMetadata), ConnectError> {
186 let addrs = resolve_target(
187 target,
188 options.enforce_dns_rebinding_check,
189 options.enforce_literal_ip_check,
190 )
191 .await?;
192 connect_to_addrs(&addrs, options.local_bind).await
193 }
194}
195
196async fn connect_to_addrs(
197 addrs: &[SocketAddr],
198 local_bind: Option<SocketAddr>,
199) -> Result<(BoxStream, ConnectionMetadata), ConnectError> {
200 let mut last_error = None;
201 for &addr in addrs {
202 let result = if let Some(local) = local_bind {
203 let local = match local {
204 SocketAddr::V6(local) => local
205 .ip()
206 .to_ipv4_mapped()
207 .map(|ip| SocketAddr::new(ip.into(), local.port()))
208 .unwrap_or(local.into()),
209 local => local,
210 };
211 let socket = if local.is_ipv4() {
212 tokio::net::TcpSocket::new_v4()
213 } else {
214 tokio::net::TcpSocket::new_v6()
215 }
216 .map_err(ConnectError::Io)?;
217 socket.bind(local).map_err(ConnectError::Io)?;
218 socket.connect(addr).await.map_err(ConnectError::Io)
219 } else {
220 TcpStream::connect(addr).await.map_err(ConnectError::Io)
221 };
222 match result {
223 Ok(stream) => {
224 let metadata = ConnectionMetadata {
225 local_addr: stream.local_addr().ok(),
226 peer_addr: stream.peer_addr().ok(),
227 };
228 return Ok((Box::new(stream), metadata));
229 }
230 Err(error) => last_error = Some(error),
231 }
232 }
233 Err(last_error.unwrap_or_else(|| ConnectError::DnsResolution("no addresses found".to_string())))
234}
235
236async fn resolve_target(
237 target: &TargetAddr,
238 enforce_dns_rebinding_check: bool,
239 enforce_literal_ip_check: bool,
240) -> Result<Vec<SocketAddr>, ConnectError> {
241 match &target.host {
242 TargetHost::Ip(ip) => {
243 if enforce_literal_ip_check && is_dns_rebinding_risk(ip) {
244 return Err(ConnectError::ReservedTarget(*ip));
245 }
246 Ok(vec![SocketAddr::new(*ip, target.port)])
247 }
248 TargetHost::Domain(domain) => {
249 let lookup = format!("{}:{}", domain, target.port);
250 let addrs: Vec<_> = tokio::net::lookup_host(&lookup)
251 .await
252 .map_err(|e| ConnectError::DnsResolution(e.to_string()))?
253 .collect();
254 if addrs.is_empty() {
255 return Err(ConnectError::DnsResolution(
256 "no addresses found".to_string(),
257 ));
258 }
259 if enforce_dns_rebinding_check {
260 if let Some(reserved) = addrs.iter().find(|addr| is_dns_rebinding_risk(&addr.ip()))
261 {
262 return Err(ConnectError::ReservedTarget(reserved.ip()));
263 }
264 }
265 Ok(addrs)
266 }
267 }
268}
269
270impl Connector for DirectConnector {
271 async fn connect(&self, target: &TargetAddr) -> Result<BoxStream, ConnectError> {
272 self.connect_with_options(target, &ConnectOptions::default())
273 .await
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use super::*;
280 use std::net::Ipv4Addr;
281 use tokio::io::{AsyncReadExt, AsyncWriteExt};
282
283 #[tokio::test]
284 async fn test_direct_connect_echo() {
285 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
286 let addr = listener.local_addr().unwrap();
287
288 let jh = tokio::spawn(async move {
289 let (mut stream, _) = listener.accept().await.unwrap();
290 let mut buf = [0u8; 1024];
291 let n = stream.read(&mut buf).await.unwrap();
292 stream.write_all(&buf[..n]).await.unwrap();
293 });
294
295 let target = TargetAddr {
296 host: TargetHost::Ip(addr.ip()),
297 port: addr.port(),
298 };
299
300 let connector = DirectConnector;
301 let mut stream = Connector::connect(&connector, &target).await.unwrap();
302
303 stream.write_all(b"ping").await.unwrap();
304 let mut buf = [0u8; 4];
305 stream.read_exact(&mut buf).await.unwrap();
306 assert_eq!(&buf, b"ping");
307
308 jh.await.unwrap();
309 }
310
311 #[tokio::test]
312 async fn direct_connect_metadata_reports_actual_socket_addresses() {
313 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
314 let peer_addr = listener.local_addr().unwrap();
315 let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
316 let target = TargetAddr {
317 host: TargetHost::Ip(peer_addr.ip()),
318 port: peer_addr.port(),
319 };
320
321 let (_stream, metadata) = DirectConnector
322 .connect_with_options_and_metadata(&target, &ConnectOptions::default())
323 .await
324 .unwrap();
325 let (server_stream, _) = accept.await.unwrap();
326 let local_addr = metadata.local_addr().expect("local socket address");
327 assert_eq!(metadata.peer_addr(), Some(peer_addr));
328 assert!(local_addr.ip().is_loopback());
329 assert_ne!(local_addr.port(), 0);
330 drop(server_stream);
331 }
332
333 #[tokio::test]
334 async fn direct_connect_metadata_uses_actual_local_bind_port() {
335 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
336 let peer_addr = listener.local_addr().unwrap();
337 let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
338 let target = TargetAddr {
339 host: TargetHost::Ip(peer_addr.ip()),
340 port: peer_addr.port(),
341 };
342
343 let (_, metadata) = DirectConnector
344 .connect_with_options_and_metadata(
345 &target,
346 &ConnectOptions {
347 local_bind: Some("127.0.0.1:0".parse().unwrap()),
348 ..Default::default()
349 },
350 )
351 .await
352 .unwrap();
353 let _ = accept.await.unwrap();
354 let local_addr = metadata.local_addr().expect("local socket address");
355 assert!(local_addr.ip().is_loopback());
356 assert_ne!(local_addr.port(), 0);
357 }
358
359 #[tokio::test]
360 async fn direct_connect_metadata_reports_ipv6_when_loopback_is_available() {
361 let listener = match tokio::net::TcpListener::bind("[::1]:0").await {
362 Ok(listener) => listener,
363 Err(error) => {
364 eprintln!("skipping IPv6 metadata check: loopback unavailable: {error}");
366 return;
367 }
368 };
369 let peer_addr = listener.local_addr().unwrap();
370 let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
371 let target = TargetAddr {
372 host: TargetHost::Ip(peer_addr.ip()),
373 port: peer_addr.port(),
374 };
375
376 let (_, metadata) = DirectConnector
377 .connect_with_options_and_metadata(&target, &ConnectOptions::default())
378 .await
379 .unwrap();
380 let _ = accept.await.unwrap();
381 assert_eq!(metadata.peer_addr(), Some(peer_addr));
382 assert!(metadata.local_addr().unwrap().ip().is_loopback());
383 }
384
385 #[tokio::test]
386 async fn dns_rebinding_policy_applies_consistently_to_domains() {
387 let target = TargetAddr {
388 host: TargetHost::Domain("localhost".to_string()),
389 port: 80,
390 };
391
392 assert!(resolve_target(&target, false, false).await.is_ok());
393 assert!(ConnectOptions::default().enforce_dns_rebinding_check);
394 assert!(matches!(
395 resolve_target(
396 &target,
397 ConnectOptions::default().enforce_dns_rebinding_check,
398 ConnectOptions::default().enforce_literal_ip_check,
399 )
400 .await,
401 Err(ConnectError::ReservedTarget(_))
402 ));
403 }
404
405 #[test]
406 fn reserved_ipv4_loopback() {
407 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
408 127, 0, 0, 1
409 ))));
410 }
411
412 #[test]
413 fn reserved_ipv4_private_10() {
414 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
415 10, 0, 0, 1
416 ))));
417 }
418
419 #[test]
420 fn reserved_ipv4_private_172() {
421 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
422 172, 16, 0, 1
423 ))));
424 }
425
426 #[test]
427 fn reserved_ipv4_private_192() {
428 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
429 192, 168, 1, 1
430 ))));
431 }
432
433 #[test]
434 fn reserved_ipv4_link_local() {
435 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
436 169, 254, 1, 1
437 ))));
438 }
439
440 #[test]
441 fn reserved_ipv4_unspecified() {
442 assert!(is_reserved_or_private_ip(&IpAddr::V4(
443 Ipv4Addr::UNSPECIFIED
444 )));
445 }
446
447 #[test]
448 fn not_reserved_ipv4_public() {
449 assert!(!is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
450 8, 8, 8, 8
451 ))));
452 }
453
454 #[test]
455 fn reserved_ipv6_loopback() {
456 assert!(is_reserved_or_private_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST)));
457 }
458
459 #[test]
460 fn reserved_ipv6_link_local() {
461 let ip = "fe80::1".parse::<Ipv6Addr>().unwrap();
462 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
463 }
464
465 #[test]
466 fn reserved_ipv4_mapped_ipv6() {
467 let ip = "::ffff:127.0.0.1".parse::<Ipv6Addr>().unwrap();
468 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
469 }
470
471 #[test]
472 fn reserved_ipv6_unique_local() {
473 let ip = "fd00::1".parse::<Ipv6Addr>().unwrap();
474 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
475 }
476
477 #[test]
478 fn reserved_ipv6_unspecified() {
479 assert!(is_reserved_or_private_ip(&IpAddr::V6(
480 Ipv6Addr::UNSPECIFIED
481 )));
482 }
483
484 #[test]
485 fn not_reserved_ipv6_public() {
486 let ip = "2606:4700:4700::1111".parse::<Ipv6Addr>().unwrap();
487 assert!(!is_reserved_or_private_ip(&IpAddr::V6(ip)));
488 }
489
490 #[test]
491 fn reserved_ipv4_multicast() {
492 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
493 224, 0, 0, 1
494 ))));
495 }
496
497 #[test]
498 fn reserved_ipv4_broadcast() {
499 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::BROADCAST)));
500 }
501
502 #[test]
503 fn reserved_ipv4_documentation() {
504 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
505 192, 0, 2, 1
506 ))));
507 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
508 198, 51, 100, 1
509 ))));
510 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
511 203, 0, 113, 1
512 ))));
513 }
514
515 #[test]
516 fn reserved_ipv4_benchmarking() {
517 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
518 198, 18, 0, 1
519 ))));
520 }
521
522 #[test]
523 fn reserved_ipv4_reserved_future() {
524 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
525 240, 0, 0, 1
526 ))));
527 }
528
529 #[test]
530 fn reserved_ipv4_this_network() {
531 assert!(is_reserved_or_private_ip(&IpAddr::V4(Ipv4Addr::new(
532 0, 1, 2, 3
533 ))));
534 }
535
536 #[test]
537 fn reserved_ipv6_multicast() {
538 let ip = "ff02::1".parse::<Ipv6Addr>().unwrap();
539 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
540 }
541
542 #[test]
543 fn reserved_ipv6_documentation() {
544 let ip = "2001:db8::1".parse::<Ipv6Addr>().unwrap();
545 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
546 }
547
548 #[test]
549 fn reserved_ipv6_discard_prefix() {
550 let ip = "0100::1".parse::<Ipv6Addr>().unwrap();
551 assert!(is_reserved_or_private_ip(&IpAddr::V6(ip)));
552 }
553
554 #[tokio::test]
555 async fn reject_domain_resolving_to_loopback() {
556 let connector = DirectConnector;
557 let target = TargetAddr {
558 host: TargetHost::Domain("localhost".to_string()),
559 port: 1,
560 };
561 let result = connector
562 .connect_with_options(
563 &target,
564 &ConnectOptions {
565 enforce_dns_rebinding_check: true,
566 ..Default::default()
567 },
568 )
569 .await;
570 assert!(matches!(result, Err(ConnectError::ReservedTarget(_))));
571 }
572
573 #[tokio::test]
574 async fn direct_connect_falls_back_to_next_resolved_address() {
575 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
576 let good_addr = listener.local_addr().unwrap();
577 let bad_addr = SocketAddr::new(good_addr.ip(), good_addr.port() + 1);
578
579 let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
580 let stream = connect_to_addrs(&[bad_addr, good_addr], None)
581 .await
582 .expect("second resolved address should be attempted");
583 drop(stream);
584 accept.await.unwrap();
585 }
586
587 #[tokio::test]
588 async fn mapped_ipv6_local_bind_uses_ipv4_socket() {
589 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
590 let addr = listener.local_addr().unwrap();
591 let accept = tokio::spawn(async move { listener.accept().await.unwrap() });
592
593 let mapped = SocketAddr::new("::ffff:127.0.0.1".parse().unwrap(), 0);
594 let stream = connect_to_addrs(&[addr], Some(mapped))
595 .await
596 .expect("mapped IPv6 local bind should connect to IPv4");
597 drop(stream);
598 accept.await.unwrap();
599 }
600}