diff options
| author | David Lönnhager <david.l@mullvad.net> | 2021-04-13 13:25:29 +0200 |
|---|---|---|
| committer | David Lönnhager <david.l@mullvad.net> | 2021-04-13 13:25:29 +0200 |
| commit | 625424d8832f1f43349659c1a5ed596c43be07e3 (patch) | |
| tree | 3a0025e89ca42b01856a442769dc0f991a141a74 | |
| parent | 0c14965d8cd012bb6c356ee85d4c505833a9023d (diff) | |
| parent | 4c9065da457998f3a707153ad6b8e0494755333b (diff) | |
| download | mullvadvpn-625424d8832f1f43349659c1a5ed596c43be07e3.tar.xz mullvadvpn-625424d8832f1f43349659c1a5ed596c43be07e3.zip | |
Merge branch 'wg-over-tcp'
| -rw-r--r-- | CHANGELOG.md | 1 | ||||
| -rw-r--r-- | Cargo.lock | 75 | ||||
| -rw-r--r-- | mullvad-cli/src/cmds/relay.rs | 96 | ||||
| -rw-r--r-- | mullvad-cli/src/format.rs | 8 | ||||
| -rw-r--r-- | mullvad-daemon/src/management_interface.rs | 35 | ||||
| -rw-r--r-- | mullvad-daemon/src/relays.rs | 1 | ||||
| -rw-r--r-- | mullvad-management-interface/proto/management_interface.proto | 1 | ||||
| -rw-r--r-- | talpid-core/Cargo.toml | 3 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/mod.rs | 6 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/wireguard/mod.rs | 99 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/connecting_state.rs | 3 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/mod.rs | 16 | ||||
| -rw-r--r-- | talpid-types/src/net/wireguard.rs | 10 |
13 files changed, 263 insertions, 91 deletions
diff --git a/CHANGELOG.md b/CHANGELOG.md index 46ef7de9ff..ee305b783e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,6 +27,7 @@ Line wrap the file at 100 chars. Th - Preserve log of old daemon instance when upgrading on Desktop. - When `MULLVAD_MANAGEMENT_SOCKET_GROUP` is set, only allow the specified group to access the management interface UDS socket. This means that only users in that group can use the CLI and GUI. +- Support WireGuard over TCP for custom VPN relays in the CLI. #### Linux - Always enable `src_valid_mark` config option when connecting to allow policty based routing. diff --git a/Cargo.lock b/Cargo.lock index f6ee1fbb81..1524765cb9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -445,8 +445,11 @@ version = "0.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "44533bbbb3bb3c1fa17d9f2e4e38bbbaf8396ba82193c4cb1b6445d711445d36" dependencies = [ + "atty", + "humantime 1.3.0", "log 0.4.14", "regex", + "termcolor", ] [[package]] @@ -456,13 +459,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f26ecb66b4bdca6c1409b40fb255eefc2bd4f6d135dab3c3124f80ffa2a9661e" dependencies = [ "atty", - "humantime", + "humantime 2.1.0", "log 0.4.14", "regex", "termcolor", ] [[package]] +name = "err-context" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "449aad22b1364e927ff3bf50f55404efd705c40065fb47f73f28704de707c89e" + +[[package]] name = "err-derive" version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -852,6 +861,15 @@ checksum = "494b4d60369511e7dea41cf646832512a94e542f68bb9c49e54518e0f468eb47" [[package]] name = "humantime" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df004cfca50ef23c36850aaaa59ad52cc70d0e90243c3c7737a4dd32dc7a3c4f" +dependencies = [ + "quick-error", +] + +[[package]] +name = "humantime" version = "2.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9a3a5bfb195931eeb336b2a7b4d761daec841b97f947d34394601737a7bba5e4" @@ -1068,9 +1086,9 @@ checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" [[package]] name = "libc" -version = "0.2.85" +version = "0.2.87" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7ccac4b00700875e6a07c6cde370d44d32fa01c5a65cdd2fca6858c479d28bb3" +checksum = "265d751d31d6780a3f956bb5b8022feba2d94eeee5a84ba64f4212eedca42213" [[package]] name = "libdbus-sys" @@ -1586,6 +1604,18 @@ dependencies = [ ] [[package]] +name = "nix" +version = "0.20.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fa9b4819da1bc61c0ea48b63b7bc8604064dd43013e7cc325df098d49cd7c18a" +dependencies = [ + "bitflags 1.2.1", + "cc", + "cfg-if 1.0.0", + "libc", +] + +[[package]] name = "notify" version = "4.0.15" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2454,6 +2484,30 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6446ced80d6c486436db5c078dde11a9f73d42b57fb273121e160b84f63d894c" [[package]] +name = "structopt" +version = "0.3.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5277acd7ee46e63e5168a80734c9f6ee81b1367a7d8772a2d765df2a3705d28c" +dependencies = [ + "clap", + "lazy_static", + "structopt-derive", +] + +[[package]] +name = "structopt-derive" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ba9cdfda491b814720b6b06e0cac513d922fc407582032e8706e9f137976f90" +dependencies = [ + "heck", + "proc-macro-error", + "proc-macro2", + "quote", + "syn", +] + +[[package]] name = "subtle" version = "2.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -2604,6 +2658,7 @@ dependencies = [ "tonic-build", "triggered", "tun", + "udp-over-tcp", "uuid", "which 4.0.2", "widestring", @@ -3134,6 +3189,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "373c8a200f9e67a0c95e62a4f52fbf80c23b4381c05a17845531982fa99e6b33" [[package]] +name = "udp-over-tcp" +version = "0.1.0" +source = "git+https://github.com/mullvad/udp-over-tcp?rev=3d1abafe112ee8c2db47ca401f8e286756454e7a#3d1abafe112ee8c2db47ca401f8e286756454e7a" +dependencies = [ + "env_logger 0.7.1", + "err-context", + "futures", + "log 0.4.14", + "nix 0.20.0", + "structopt", + "tokio", +] + +[[package]] name = "unicode-segmentation" version = "1.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" diff --git a/mullvad-cli/src/cmds/relay.rs b/mullvad-cli/src/cmds/relay.rs index 6343b9e5d5..6aa2ea5f2c 100644 --- a/mullvad-cli/src/cmds/relay.rs +++ b/mullvad-cli/src/cmds/relay.rs @@ -44,72 +44,71 @@ impl Command for Relay { .arg( clap::Arg::with_name("host") .help("Hostname or IP") - .required(true) - .index(1), + .required(true), ) .arg( clap::Arg::with_name("port") .help("Remote network port") - .required(true) - .index(2), + .required(true), ) .arg( - clap::Arg::with_name("peer-key") + clap::Arg::with_name("peer-pubkey") .help("Base64 encoded peer public key") - .index(3) - .required(false), + .required(true), ) .arg( clap::Arg::with_name("v4-gateway") .help("IPv4 gateway address") - .long("v4-gateway") - .index(4) - .required(false), - ).arg( - clap::Arg::with_name("v6-gateway") - .help("IPv6 gateway address") - .long("v6-gateway") - .takes_value(true) - .required(false), + .required(true), ) .arg( clap::Arg::with_name("addr") .help("Local address of wireguard tunnel") - .long("addr") - .takes_value(true) - .multiple(true) - .required(false), - ), + .required(true) + .multiple(true), + ) + .arg( + clap::Arg::with_name("protocol") + .help("Transport protocol. If TCP is selected, traffic is \ + sent over TCP using a udp-over-tcp proxy") + .long("protocol") + .default_value("udp") + .possible_values(&["udp", "tcp"]), + ) + .arg( + clap::Arg::with_name("v6-gateway") + .help("IPv6 gateway address") + .long("v6-gateway") + .takes_value(true), + ) ) .subcommand(clap::SubCommand::with_name("openvpn") .arg( clap::Arg::with_name("host") .help("Hostname or IP") - .required(true) - .index(1), + .required(true), ) .arg( clap::Arg::with_name("port") .help("Remote network port") - .required(true) - .index(2), - ) - .arg( - clap::Arg::with_name("protocol") - .help("Transport protocol. For Wireguard this is ignored.") - .index(3) - .default_value("udp") - .possible_values(&["udp", "tcp"]), + .required(true), ) .arg( clap::Arg::with_name("username") .help("Username to be used with the OpenVpn relay") - .index(4), + .required(true), ) .arg( clap::Arg::with_name("password") .help("Password to be used with the OpenVpn relay") - .index(5), + .required(true), + ) + .arg( + clap::Arg::with_name("protocol") + .help("Transport protocol") + .long("protocol") + .default_value("udp") + .possible_values(&["udp", "tcp"]), ) ) ) @@ -152,14 +151,12 @@ impl Command for Relay { .arg( clap::Arg::with_name("transport protocol") .long("protocol") - .required(false) .default_value("any") .possible_values(&["any", "udp", "tcp"]), ) .arg( clap::Arg::with_name("ip version") .long("ipv") - .required(false) .default_value("any") .possible_values(&["any", "4", "6"]), ), @@ -248,15 +245,7 @@ impl Relay { let password = value_t!(matches.value_of("password"), String).unwrap_or_else(|e| e.exit()); let protocol = value_t!(matches.value_of("protocol"), String).unwrap_or_else(|e| e.exit()); - let protocol = match protocol.as_str() { - "udp" => TransportProtocol::Udp, - "tcp" => TransportProtocol::Tcp, - _ => clap::Error::with_description( - "unknown transport protocol", - clap::ErrorKind::ValueValidation, - ) - .exit(), - }; + let protocol = Self::validate_transport_protocol(&protocol); CustomRelaySettings { host, @@ -278,7 +267,7 @@ impl Relay { let port = value_t!(matches.value_of("port"), u16).unwrap_or_else(|e| e.exit()); let addresses = values_t!(matches.values_of("addr"), IpAddr).unwrap_or_else(|e| e.exit()); let peer_key_str = - value_t!(matches.value_of("peer-key"), String).unwrap_or_else(|e| e.exit()); + value_t!(matches.value_of("peer-pubkey"), String).unwrap_or_else(|e| e.exit()); let ipv4_gateway = value_t!(matches.value_of("v4-gateway"), Ipv4Addr).unwrap_or_else(|e| e.exit()); let ipv6_gateway = match value_t!(matches.value_of("v6-gateway"), Ipv6Addr) { @@ -288,6 +277,8 @@ impl Relay { _ => e.exit(), }, }; + let protocol = value_t!(matches.value_of("protocol"), String).unwrap_or_else(|e| e.exit()); + let protocol = Self::validate_transport_protocol(&protocol); let mut private_key_str = String::new(); println!("Reading private key from standard input"); let _ = io::stdin().lock().read_line(&mut private_key_str); @@ -316,6 +307,7 @@ impl Relay { .collect(), endpoint: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), port) .to_string(), + protocol: protocol as i32, }), ipv4_gateway: ipv4_gateway.to_string(), ipv6_gateway: ipv6_gateway @@ -346,6 +338,18 @@ impl Relay { key } + fn validate_transport_protocol(protocol: &str) -> TransportProtocol { + match protocol { + "udp" => TransportProtocol::Udp, + "tcp" => TransportProtocol::Tcp, + _ => clap::Error::with_description( + "invalid transport protocol", + clap::ErrorKind::ValueValidation, + ) + .exit(), + } + } + async fn set_hostname(&self, matches: &clap::ArgMatches<'_>) -> Result<()> { let hostname = matches.value_of("hostname").unwrap(); let countries = Self::get_filtered_relays().await?; diff --git a/mullvad-cli/src/format.rs b/mullvad-cli/src/format.rs index 7070e322ac..250bc70ba7 100644 --- a/mullvad-cli/src/format.rs +++ b/mullvad-cli/src/format.rs @@ -59,13 +59,13 @@ pub fn print_state(state: &TunnelState) { fn format_endpoint(endpoint: &TunnelEndpoint) -> String { let mut out = format!( "{} {} over {}", - match TunnelType::from_i32(endpoint.tunnel_type).expect("unknown tunnel protocol") { + match TunnelType::from_i32(endpoint.tunnel_type).expect("invalid tunnel protocol") { TunnelType::Wireguard => "WireGuard", TunnelType::Openvpn => "OpenVPN", }, endpoint.address, format_protocol( - TransportProtocol::from_i32(endpoint.protocol).expect("unknown transport protocol") + TransportProtocol::from_i32(endpoint.protocol).expect("invalid transport protocol") ), ); @@ -73,13 +73,13 @@ fn format_endpoint(endpoint: &TunnelEndpoint) -> String { write!( &mut out, " via {} {} over {}", - match ProxyType::from_i32(proxy.proxy_type).expect("unknown proxy type") { + match ProxyType::from_i32(proxy.proxy_type).expect("invalid proxy type") { ProxyType::Shadowsocks => "Shadowsocks", ProxyType::Custom => "custom bridge", }, proxy.address, format_protocol( - TransportProtocol::from_i32(proxy.protocol).expect("unknown transport protocol") + TransportProtocol::from_i32(proxy.protocol).expect("invalid transport protocol") ), ) .unwrap(); diff --git a/mullvad-daemon/src/management_interface.rs b/mullvad-daemon/src/management_interface.rs index dff318cece..05bd38e1dc 100644 --- a/mullvad-daemon/src/management_interface.rs +++ b/mullvad-daemon/src/management_interface.rs @@ -808,15 +808,7 @@ fn convert_relay_settings_update( ConnectionConfig::OpenVpn(openvpn::ConnectionConfig { endpoint: net::Endpoint { address, - protocol: match types::TransportProtocol::from_i32(config.protocol) { - Some(types::TransportProtocol::Udp) => TransportProtocol::Udp, - Some(types::TransportProtocol::Tcp) => TransportProtocol::Tcp, - None => { - return Err(Status::invalid_argument( - "invalid transport protocol", - )) - } - }, + protocol: convert_proto_transport_protocol(config.protocol)?, }, username: config.username.clone(), password: config.password.clone(), @@ -893,6 +885,7 @@ fn convert_relay_settings_update( public_key: wireguard::PublicKey::from(public_key), allowed_ips, endpoint, + protocol: convert_proto_transport_protocol(peer.protocol)?, }, ipv4_gateway, ipv6_gateway, @@ -923,7 +916,7 @@ fn convert_relay_settings_update( Some(types::TunnelType::Wireguard) => { Some(Constraint::Only(TunnelType::Wireguard)) } - None => return Err(Status::invalid_argument("unknown tunnel protocol")), + None => return Err(Status::invalid_argument("invalid tunnel protocol")), }, None => Some(Constraint::Any), } @@ -934,13 +927,7 @@ fn convert_relay_settings_update( let transport_protocol = if let Some(ref constraints) = settings.openvpn_constraints { match &constraints.protocol { Some(constraint) => { - match types::TransportProtocol::from_i32(constraint.protocol) { - Some(types::TransportProtocol::Udp) => Some(TransportProtocol::Udp), - Some(types::TransportProtocol::Tcp) => Some(TransportProtocol::Tcp), - None => { - return Err(Status::invalid_argument("unknown transport protocol")) - } - } + Some(convert_proto_transport_protocol(constraint.protocol)?) } None => None, } @@ -967,7 +954,7 @@ fn convert_relay_settings_update( Some(types::IpVersion::V4) => Some(IpVersion::V4), Some(types::IpVersion::V6) => Some(IpVersion::V6), None => { - return Err(Status::invalid_argument("unknown ip protocol version")) + return Err(Status::invalid_argument("invalid ip protocol version")) } }, None => None, @@ -1103,6 +1090,10 @@ fn convert_connection_config(config: &ConnectionConfig) -> types::ConnectionConf .map(|address| address.to_string()) .collect(), endpoint: config.peer.endpoint.to_string(), + protocol: i32::from(match config.peer.protocol { + TransportProtocol::Udp => types::TransportProtocol::Udp, + TransportProtocol::Tcp => types::TransportProtocol::Tcp, + }), }), ipv4_gateway: config.ipv4_gateway.to_string(), ipv6_gateway: config @@ -1546,6 +1537,14 @@ fn convert_proto_location(location: types::RelayLocation) -> Constraint<Location } } +fn convert_proto_transport_protocol(protocol: i32) -> Result<TransportProtocol, Status> { + match types::TransportProtocol::from_i32(protocol) { + Some(types::TransportProtocol::Udp) => Ok(TransportProtocol::Udp), + Some(types::TransportProtocol::Tcp) => Ok(TransportProtocol::Tcp), + None => Err(Status::invalid_argument("invalid transport protocol")), + } +} + pub struct ManagementInterfaceServer { subscriptions: Arc<RwLock<Vec<EventsListenerSender>>>, socket_path: String, diff --git a/mullvad-daemon/src/relays.rs b/mullvad-daemon/src/relays.rs index 54d2074ad9..c029baaa54 100644 --- a/mullvad-daemon/src/relays.rs +++ b/mullvad-daemon/src/relays.rs @@ -697,6 +697,7 @@ impl RelaySelector { public_key: data.public_key, endpoint: SocketAddr::new(host, port), allowed_ips: all_of_the_internet(), + protocol: TransportProtocol::Udp, }; Some(MullvadEndpoint::Wireguard { peer: peer_config, diff --git a/mullvad-management-interface/proto/management_interface.proto b/mullvad-management-interface/proto/management_interface.proto index 9404baaaf3..ba36deeea4 100644 --- a/mullvad-management-interface/proto/management_interface.proto +++ b/mullvad-management-interface/proto/management_interface.proto @@ -340,6 +340,7 @@ message ConnectionConfig { bytes public_key = 1; repeated string allowed_ips = 2; string endpoint = 3; + TransportProtocol protocol = 4; } TunnelConfig tunnel = 1; diff --git a/talpid-core/Cargo.toml b/talpid-core/Cargo.toml index 7c01358484..5413634702 100644 --- a/talpid-core/Cargo.toml +++ b/talpid-core/Cargo.toml @@ -26,8 +26,9 @@ talpid-types = { path = "../talpid-types" } uuid = { version = "0.8", features = ["v4"] } zeroize = "1" chrono = "0.4" -tokio = { version = "0.2", features = [ "process", "rt-threaded", "stream" ] } +tokio = { version = "0.2", features = [ "process", "rt-threaded", "stream" ] } rand = "0.7" +udp-over-tcp = { git = "https://github.com/mullvad/udp-over-tcp", rev = "3d1abafe112ee8c2db47ca401f8e286756454e7a" } [target.'cfg(not(target_os="android"))'.dependencies] diff --git a/talpid-core/src/tunnel/mod.rs b/talpid-core/src/tunnel/mod.rs index a5b763cad1..87fc3d4306 100644 --- a/talpid-core/src/tunnel/mod.rs +++ b/talpid-core/src/tunnel/mod.rs @@ -148,6 +148,7 @@ impl TunnelMonitor { /// on tunnel state changes. #[cfg_attr(any(target_os = "android", windows), allow(unused_variables))] pub fn start<L>( + runtime: tokio::runtime::Handle, tunnel_parameters: &TunnelParameters, log_dir: &Option<PathBuf>, resource_dir: &Path, @@ -170,6 +171,7 @@ impl TunnelMonitor { TunnelParameters::OpenVpn(_) => Err(Error::UnsupportedPlatform), TunnelParameters::Wireguard(config) => Self::start_wireguard_tunnel( + runtime, &config, log_file, on_event, @@ -200,6 +202,7 @@ impl TunnelMonitor { } fn start_wireguard_tunnel<L>( + runtime: tokio::runtime::Handle, params: &wireguard_types::TunnelParameters, log: Option<PathBuf>, on_event: L, @@ -211,7 +214,8 @@ impl TunnelMonitor { { let config = wireguard::config::Config::from_parameters(¶ms)?; let monitor = wireguard::WireguardMonitor::start( - &config, + runtime, + config, log.as_ref().map(|p| p.as_path()), on_event, tun_provider, diff --git a/talpid-core/src/tunnel/wireguard/mod.rs b/talpid-core/src/tunnel/wireguard/mod.rs index 9fe6bec4e7..81089e59e6 100644 --- a/talpid-core/src/tunnel/wireguard/mod.rs +++ b/talpid-core/src/tunnel/wireguard/mod.rs @@ -3,16 +3,19 @@ use self::config::Config; use super::tun_provider; use super::{tun_provider::TunProvider, TunnelEvent, TunnelMetadata}; use crate::routing::{self, RequiredRoute}; +use futures::future::abortable; #[cfg(target_os = "linux")] use lazy_static::lazy_static; #[cfg(target_os = "linux")] use std::env; use std::{ collections::HashSet, + net::SocketAddr, path::Path, sync::{mpsc, Arc, Mutex}, }; -use talpid_types::ErrorExt; +use talpid_types::{net::TransportProtocol, ErrorExt}; +use udp_over_tcp::{TcpOptions, Udp2Tcp}; /// WireGuard config data-types pub mod config; @@ -29,6 +32,7 @@ type Result<T> = std::result::Result<T, Error>; /// Errors that can happen in the Wireguard tunnel monitor. #[derive(err_derive::Error, Debug)] +#[error(no_from)] pub enum Error { /// Failed to set up routing. #[error(display = "Failed to setup routing")] @@ -42,6 +46,14 @@ pub enum Error { #[error(display = "Tunnel failed")] TunnelError(#[error(source)] TunnelError), + /// Failed to set up Udp2Tcp + #[error(display = "Failed to start UDP-over-TCP proxy")] + Udp2TcpError(#[error(source)] udp_over_tcp::udp2tcp::ConnectError), + + /// Failed to obtain the local UDP socket address + #[error(display = "Failed obtain local address for the UDP socket in Udp2Tcp")] + GetLocalUdpAddress(#[error(source)] std::io::Error), + /// Failed to setup connectivity monitor #[error(display = "Connectivity monitor failed")] ConnectivityMonitorError(#[error(source)] connectivity_check::Error), @@ -57,6 +69,7 @@ pub struct WireguardMonitor { close_msg_sender: mpsc::Sender<CloseMsg>, close_msg_receiver: mpsc::Receiver<CloseMsg>, pinger_stop_sender: mpsc::Sender<()>, + _tcp_proxies: Vec<TcpProxy>, } #[cfg(target_os = "linux")] @@ -71,15 +84,78 @@ lazy_static! { .unwrap_or(false); } +struct TcpProxy { + local_addr: SocketAddr, + abort_handle: futures::future::AbortHandle, +} + +impl TcpProxy { + pub fn new(runtime: &tokio::runtime::Handle, endpoint: SocketAddr) -> Result<Self> { + let listen_addr = if endpoint.is_ipv4() { + SocketAddr::new("127.0.0.1".parse().unwrap(), 0) + } else { + SocketAddr::new("::1".parse().unwrap(), 0) + }; + + let udp2tcp = runtime + .block_on(Udp2Tcp::new( + listen_addr, + endpoint, + Some(&TcpOptions { + #[cfg(target_os = "linux")] + fwmark: Some(crate::linux::TUNNEL_FW_MARK), + ..TcpOptions::default() + }), + )) + .map_err(Error::Udp2TcpError)?; + + let local_addr = udp2tcp + .local_udp_addr() + .map_err(Error::GetLocalUdpAddress)?; + + let (udp2tcp_future, abort_handle) = abortable(udp2tcp.run()); + runtime.spawn(udp2tcp_future); + + Ok(Self { + local_addr, + abort_handle, + }) + } + + pub fn local_udp_addr(&self) -> SocketAddr { + self.local_addr + } +} + +impl Drop for TcpProxy { + fn drop(&mut self) { + self.abort_handle.abort(); + } +} + impl WireguardMonitor { /// Starts a WireGuard tunnel with the given config pub fn start<F: Fn(TunnelEvent) + Send + Sync + Clone + 'static>( - config: &Config, + runtime: tokio::runtime::Handle, + mut config: Config, log_path: Option<&Path>, on_event: F, tun_provider: &mut TunProvider, route_manager: &mut routing::RouteManager, ) -> Result<WireguardMonitor> { + let mut tcp_proxies = vec![]; + + for peer in &mut config.peers { + if peer.protocol == TransportProtocol::Tcp { + let udp2tcp = TcpProxy::new(&runtime, peer.endpoint.clone())?; + + // Replace remote peer with proxy + peer.endpoint = udp2tcp.local_udp_addr(); + + tcp_proxies.push(udp2tcp); + } + } + let tunnel = Self::open_tunnel(&config, log_path, tun_provider, route_manager)?; let iface_name = tunnel.get_interface_name().to_string(); @@ -105,6 +181,7 @@ impl WireguardMonitor { close_msg_sender, close_msg_receiver, pinger_stop_sender: pinger_tx, + _tcp_proxies: tcp_proxies, }; let metadata = Self::tunnel_metadata(&iface_name, &config); @@ -115,7 +192,8 @@ impl WireguardMonitor { iface_name.to_string(), Arc::downgrade(&monitor.tunnel), pinger_rx, - )?; + ) + .map_err(Error::ConnectivityMonitorError)?; std::thread::spawn(move || { match connectivity_monitor.establish_connectivity() { @@ -188,12 +266,15 @@ impl WireguardMonitor { #[cfg(target_os = "linux")] log::debug!("Using userspace WireGuard implementation"); - Ok(Box::new(WgGoTunnel::start_tunnel( - &config, - log_path, - tun_provider, - Self::get_tunnel_routes(config), - )?)) + Ok(Box::new( + WgGoTunnel::start_tunnel( + &config, + log_path, + tun_provider, + Self::get_tunnel_routes(config), + ) + .map_err(Error::TunnelError)?, + )) } /// Returns a close handle for the tunnel diff --git a/talpid-core/src/tunnel_state_machine/connecting_state.rs b/talpid-core/src/tunnel_state_machine/connecting_state.rs index f9f4a00764..5257475d63 100644 --- a/talpid-core/src/tunnel_state_machine/connecting_state.rs +++ b/talpid-core/src/tunnel_state_machine/connecting_state.rs @@ -89,6 +89,7 @@ impl ConnectingState { } fn start_tunnel( + runtime: tokio::runtime::Handle, parameters: TunnelParameters, log_dir: &Option<PathBuf>, resource_dir: &Path, @@ -102,6 +103,7 @@ impl ConnectingState { }; let monitor = TunnelMonitor::start( + runtime, ¶meters, log_dir, resource_dir, @@ -420,6 +422,7 @@ impl TunnelState for ConnectingState { } match Self::start_tunnel( + shared_values.runtime.clone(), tunnel_parameters, &shared_values.log_dir, &shared_values.resource_dir, diff --git a/talpid-core/src/tunnel_state_machine/mod.rs b/talpid-core/src/tunnel_state_machine/mod.rs index 90012a8296..509a1a0e43 100644 --- a/talpid-core/src/tunnel_state_machine/mod.rs +++ b/talpid-core/src/tunnel_state_machine/mod.rs @@ -139,7 +139,7 @@ pub async fn spawn( } }; - state_machine.run(runtime, state_change_listener); + state_machine.run(state_change_listener); if shutdown_tx.send(()).is_err() { log::error!("Can't send shutdown completion to daemon"); @@ -222,9 +222,10 @@ impl TunnelStateMachine { let firewall = Firewall::new(args).map_err(Error::InitFirewallError)?; let dns_monitor = DnsMonitor::new(cache_dir).map_err(Error::InitDnsMonitorError)?; - let route_manager = - RouteManager::new(runtime, HashSet::new()).map_err(Error::InitRouteManagerError)?; + let route_manager = RouteManager::new(runtime.clone(), HashSet::new()) + .map_err(Error::InitRouteManagerError)?; let mut shared_values = SharedTunnelStateValues { + runtime, firewall, dns_monitor, route_manager, @@ -250,13 +251,11 @@ impl TunnelStateMachine { }) } - fn run( - mut self, - runtime: tokio::runtime::Handle, - change_listener: impl Sender<TunnelStateTransition> + Send + 'static, - ) { + fn run(mut self, change_listener: impl Sender<TunnelStateTransition> + Send + 'static) { use EventConsequence::*; + let runtime = self.shared_values.runtime.clone(); + while let Some(state_wrapper) = self.current_state.take() { match state_wrapper.handle_event(&runtime, &mut self.commands, &mut self.shared_values) { @@ -295,6 +294,7 @@ pub trait TunnelParametersGenerator: Send + 'static { /// Values that are common to all tunnel states. struct SharedTunnelStateValues { + runtime: tokio::runtime::Handle, firewall: Firewall, dns_monitor: DnsMonitor, route_manager: RouteManager, diff --git a/talpid-types/src/net/wireguard.rs b/talpid-types/src/net/wireguard.rs index 731d93c79e..60a0a29a6d 100644 --- a/talpid-types/src/net/wireguard.rs +++ b/talpid-types/src/net/wireguard.rs @@ -34,7 +34,7 @@ impl ConnectionConfig { pub fn get_endpoint(&self) -> Endpoint { Endpoint { address: self.peer.endpoint, - protocol: TransportProtocol::Udp, + protocol: self.peer.protocol, } } } @@ -47,6 +47,14 @@ pub struct PeerConfig { pub allowed_ips: Vec<IpNetwork>, /// IP address of the WireGuard server. pub endpoint: SocketAddr, + /// Transport protocol. WireGuard only supports UDP directly. + /// If this is set to TCP, then traffic is proxied using [`udp_to_tcp::Udp2Tcp`]. + #[serde(default = "default_peer_transport")] + pub protocol: TransportProtocol, +} + +fn default_peer_transport() -> TransportProtocol { + TransportProtocol::Udp } #[derive(Clone, Eq, PartialEq, Deserialize, Serialize, Debug)] |
