|
1 | 1 | use clap::Parser; |
2 | 2 | use ipnet::IpNet; |
3 | 3 | use rand::Rng; |
4 | | -use socket2::{Domain, Protocol, Socket, Type}; |
5 | | -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; |
| 4 | +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, ToSocketAddrs}; |
6 | 5 | use std::sync::Arc; |
7 | 6 | use thiserror::Error; |
8 | 7 | use tokio::io::{AsyncReadExt, AsyncWriteExt}; |
9 | 8 | use tokio::net::{TcpListener, TcpSocket, TcpStream}; |
10 | 9 | use tracing::{error, info, warn}; |
11 | 10 |
|
12 | | -#[cfg(unix)] |
13 | | -use std::os::unix::io::{FromRawFd, IntoRawFd}; |
14 | | - |
15 | 11 | // SOCKS5 protocol constants |
16 | 12 | const SOCKS_VERSION: u8 = 0x05; |
17 | 13 |
|
@@ -342,39 +338,40 @@ async fn handle_request(stream: &mut TcpStream, config: &ServerConfig) -> Result |
342 | 338 | let local_ip = random_ip_from_cidr(cidr); |
343 | 339 | info!("Connecting to {} via {}", target, local_ip); |
344 | 340 |
|
345 | | - // Resolve target address (handles both IP and domain names) |
346 | | - let mut addrs = tokio::net::lookup_host(&target).await?; |
347 | | - let remote_addr = match local_ip { |
348 | | - IpAddr::V4(_) => addrs.find(|a| a.is_ipv4()), |
349 | | - IpAddr::V6(_) => addrs.find(|a| a.is_ipv6()), |
350 | | - }; |
351 | | - |
352 | | - match remote_addr { |
353 | | - Some(addr) => { |
354 | | - let domain = match local_ip { |
355 | | - IpAddr::V4(_) => Domain::IPV4, |
356 | | - IpAddr::V6(_) => Domain::IPV6, |
357 | | - }; |
358 | | - let socket = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?; |
359 | | - |
360 | | - // Enable IP_FREEBIND to bind to addresses not yet configured on the interface |
361 | | - #[cfg(target_os = "linux")] |
362 | | - socket.set_freebind(true)?; |
363 | | - |
364 | | - socket.bind(&SocketAddr::new(local_ip, 0).into())?; |
365 | | - socket.set_nonblocking(true)?; |
366 | | - |
367 | | - // Convert socket2::Socket to tokio::TcpSocket via raw fd |
368 | | - #[cfg(unix)] |
369 | | - let tcp_socket = unsafe { TcpSocket::from_raw_fd(socket.into_raw_fd()) }; |
370 | | - |
371 | | - tcp_socket.connect(addr).await |
372 | | - } |
373 | | - None => { |
374 | | - // Fallback: try to connect without binding if no matching address family |
375 | | - warn!("No matching address family for {}, connecting without bind", target); |
376 | | - TcpStream::connect(&target).await |
| 341 | + // Resolve target address and try to connect |
| 342 | + match target.to_socket_addrs() { |
| 343 | + Ok(addrs) => { |
| 344 | + let mut last_err = None; |
| 345 | + let mut connected = None; |
| 346 | + |
| 347 | + for addr in addrs { |
| 348 | + // Create socket matching the local IP family |
| 349 | + let socket = match local_ip { |
| 350 | + IpAddr::V4(_) if addr.is_ipv4() => TcpSocket::new_v4()?, |
| 351 | + IpAddr::V6(_) if addr.is_ipv6() => TcpSocket::new_v6()?, |
| 352 | + _ => continue, // Skip if address family doesn't match |
| 353 | + }; |
| 354 | + |
| 355 | + let bind_addr = SocketAddr::new(local_ip, 0); |
| 356 | + if socket.bind(bind_addr).is_ok() { |
| 357 | + match socket.connect(addr).await { |
| 358 | + Ok(stream) => { |
| 359 | + connected = Some(stream); |
| 360 | + break; |
| 361 | + } |
| 362 | + Err(e) => last_err = Some(e), |
| 363 | + } |
| 364 | + } |
| 365 | + } |
| 366 | + |
| 367 | + match connected { |
| 368 | + Some(stream) => Ok(stream), |
| 369 | + None => Err(last_err.unwrap_or_else(|| { |
| 370 | + std::io::Error::new(std::io::ErrorKind::Other, "No matching address family") |
| 371 | + })), |
| 372 | + } |
377 | 373 | } |
| 374 | + Err(e) => Err(e), |
378 | 375 | } |
379 | 376 | } else { |
380 | 377 | info!("Connecting to {}", target); |
|
0 commit comments