diff options
| author | Janito Vaqueiro Ferreira Filho <janito@mullvad.net> | 2020-02-06 09:28:32 -0300 |
|---|---|---|
| committer | Janito Vaqueiro Ferreira Filho <janito@mullvad.net> | 2020-02-06 09:28:32 -0300 |
| commit | 8ced6b2a795352e99d212292c17378a7885b943a (patch) | |
| tree | 0a89a8d3076d0da44de31ec84e1adfe085aa16c7 | |
| parent | 982b11d0084a205e1305c621ce430bb5f1040b6d (diff) | |
| parent | 00701395ea0625789ad915c532a4c2417a6985b0 (diff) | |
| download | mullvadvpn-8ced6b2a795352e99d212292c17378a7885b943a.tar.xz mullvadvpn-8ced6b2a795352e99d212292c17378a7885b943a.zip | |
Merge branch 'allow-lan-on-android'
| -rw-r--r-- | talpid-core/src/firewall/mod.rs | 4 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/tun_provider/android/ipnetwork_sub.rs | 452 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/tun_provider/android/mod.rs (renamed from talpid-core/src/tunnel/tun_provider/android.rs) | 67 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/tun_provider/mod.rs | 7 | ||||
| -rw-r--r-- | talpid-core/src/tunnel/wireguard/wireguard_go.rs | 14 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/connected_state.rs | 30 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/connecting_state.rs | 107 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/disconnected_state.rs | 7 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/disconnecting_state.rs | 6 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/error_state.rs | 9 | ||||
| -rw-r--r-- | talpid-core/src/tunnel_state_machine/mod.rs | 26 |
11 files changed, 640 insertions, 89 deletions
diff --git a/talpid-core/src/firewall/mod.rs b/talpid-core/src/firewall/mod.rs index 93279437a1..9c81dd63df 100644 --- a/talpid-core/src/firewall/mod.rs +++ b/talpid-core/src/firewall/mod.rs @@ -31,7 +31,7 @@ pub use self::imp::Error; #[cfg(unix)] lazy_static! { /// When "allow local network" is enabled the app will allow traffic to and from these networks. - static ref ALLOWED_LAN_NETS: [IpNetwork; 5] = [ + pub(crate) static ref ALLOWED_LAN_NETS: [IpNetwork; 5] = [ IpNetwork::V4(Ipv4Network::new(Ipv4Addr::new(10, 0, 0, 0), 8).unwrap()), IpNetwork::V4(Ipv4Network::new(Ipv4Addr::new(172, 16, 0, 0), 12).unwrap()), IpNetwork::V4(Ipv4Network::new(Ipv4Addr::new(192, 168, 0, 0), 16).unwrap()), @@ -39,7 +39,7 @@ lazy_static! { IpNetwork::V6(Ipv6Network::new(Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 0), 10).unwrap()), ]; /// When "allow local network" is enabled the app will allow traffic to these networks. - static ref ALLOWED_LAN_MULTICAST_NETS: [IpNetwork; 5] = [ + pub(crate) static ref ALLOWED_LAN_MULTICAST_NETS: [IpNetwork; 5] = [ // Local subnetwork multicast. Not routable IpNetwork::V4(Ipv4Network::new(Ipv4Addr::new(224, 0, 0, 0), 24).unwrap()), // Simple Service Discovery Protocol (SSDP) address diff --git a/talpid-core/src/tunnel/tun_provider/android/ipnetwork_sub.rs b/talpid-core/src/tunnel/tun_provider/android/ipnetwork_sub.rs new file mode 100644 index 0000000000..ffaa585302 --- /dev/null +++ b/talpid-core/src/tunnel/tun_provider/android/ipnetwork_sub.rs @@ -0,0 +1,452 @@ +use ipnetwork::{IpNetwork, Ipv4Network, Ipv6Network}; +use std::{ + fmt::Debug, + iter, + marker::PhantomData, + ops::{Add, BitAnd, BitXor, Not, Shl, Sub}, +}; + +pub trait AbstractIpNetwork: Clone + Copy + 'static { + type Representation: Add<Output = Self::Representation> + + BitAnd<Output = Self::Representation> + + BitXor<Output = Self::Representation> + + Clone + + Copy + + Debug + + PartialEq + + Not<Output = Self::Representation> + + Shl<u8, Output = Self::Representation> + + Sub<Output = Self::Representation>; + + const ZERO: Self::Representation; + const ONE: Self::Representation; + const MAX_PREFIX: u8; + + fn new(network: Self::Representation, prefix: u8) -> Self; + fn mask(self) -> Self::Representation; + fn network(self) -> Self::Representation; + fn prefix(self) -> u8; +} + +impl AbstractIpNetwork for Ipv4Network { + type Representation = u32; + + const ZERO: Self::Representation = 0; + const ONE: Self::Representation = 1; + const MAX_PREFIX: u8 = 32; + + fn new(network: Self::Representation, prefix: u8) -> Self { + Ipv4Network::new(network.into(), prefix).expect("Invalid IPv4 network prefix") + } + + fn mask(self) -> Self::Representation { + Ipv4Network::mask(&self).into() + } + + fn network(self) -> Self::Representation { + Ipv4Network::network(&self).into() + } + + fn prefix(self) -> u8 { + Ipv4Network::prefix(&self) + } +} + +impl AbstractIpNetwork for Ipv6Network { + type Representation = u128; + + const ZERO: Self::Representation = 0; + const ONE: Self::Representation = 1; + const MAX_PREFIX: u8 = 128; + + fn new(network: Self::Representation, prefix: u8) -> Self { + Ipv6Network::new(network.into(), prefix).expect("Invalid IPv6 network prefix") + } + + fn mask(self) -> Self::Representation { + Ipv6Network::mask(&self).into() + } + + fn network(self) -> Self::Representation { + Ipv6Network::network(&self).into() + } + + fn prefix(self) -> u8 { + Ipv6Network::prefix(&self) + } +} + +#[derive(Clone, Copy, Debug)] +pub struct IpNetworkRange<T: AbstractIpNetwork> { + network: T::Representation, + bit_position: u8, + max_bit_position: u8, + _network_type: PhantomData<T>, +} + +impl<T> Iterator for IpNetworkRange<T> +where + T: AbstractIpNetwork, +{ + type Item = T; + + fn next(&mut self) -> Option<Self::Item> { + if self.bit_position < self.max_bit_position { + let bit_mask = T::ONE << self.bit_position; + let prefix_mask = !(bit_mask - T::ONE); + let address = (self.network ^ bit_mask) & prefix_mask; + let prefix = T::MAX_PREFIX - self.bit_position; + + self.bit_position += 1; + + Some(T::new(address, prefix)) + } else { + None + } + } +} + +#[derive(Clone, Copy, Debug)] +pub enum IpNetworks<T: AbstractIpNetwork> { + Empty, + SingleNetwork(T), + MultipleNetworks(IpNetworkRange<T>), +} + +impl<T> Iterator for IpNetworks<T> +where + T: AbstractIpNetwork, +{ + type Item = T; + + fn next(&mut self) -> Option<Self::Item> { + match self { + IpNetworks::Empty => None, + &mut IpNetworks::SingleNetwork(network) => { + *self = IpNetworks::Empty; + Some(network) + } + IpNetworks::MultipleNetworks(range) => { + if let Some(item) = range.next() { + Some(item) + } else { + *self = IpNetworks::Empty; + None + } + } + } + } +} + +pub trait IpNetworkSub: Copy + Sized + 'static { + type Output: Iterator<Item = Self>; + + fn sub(self, other: Self) -> Self::Output; + + fn sub_all(self, others: impl IntoIterator<Item = Self>) -> Box<dyn Iterator<Item = Self>> { + let mut result: Box<dyn Iterator<Item = Self>> = Box::new(iter::once(self)); + + for other in others { + result = Box::new(result.flat_map(move |network| network.sub(other))); + } + + result + } +} + +impl<T> IpNetworkSub for T +where + T: AbstractIpNetwork, +{ + type Output = IpNetworks<T>; + + fn sub(self, other: Self) -> Self::Output { + let subtrahend = self.network(); + let minuend = other.network(); + let mask = self.mask(); + + if minuend & mask == subtrahend { + let max_bit_position = T::MAX_PREFIX - self.prefix(); + let bit_position = T::MAX_PREFIX - other.prefix(); + + IpNetworks::MultipleNetworks(IpNetworkRange { + network: minuend, + bit_position, + max_bit_position, + _network_type: PhantomData, + }) + } else { + let other_mask = other.mask(); + + if subtrahend & other_mask == minuend { + IpNetworks::Empty + } else { + IpNetworks::SingleNetwork(self) + } + } + } +} + +#[derive(Debug)] +pub enum IpNetworkIterator { + V4(IpNetworks<Ipv4Network>), + V6(IpNetworks<Ipv6Network>), +} + +impl Iterator for IpNetworkIterator { + type Item = IpNetwork; + + fn next(&mut self) -> Option<Self::Item> { + match self { + IpNetworkIterator::V4(iterator) => iterator.next().map(IpNetwork::V4), + IpNetworkIterator::V6(iterator) => iterator.next().map(IpNetwork::V6), + } + } +} + +impl IpNetworkSub for IpNetwork { + type Output = IpNetworkIterator; + + fn sub(self, other: Self) -> Self::Output { + match (self, other) { + (IpNetwork::V4(self_v4), IpNetwork::V4(other_v4)) => { + IpNetworkIterator::V4(self_v4.sub(other_v4)) + } + (IpNetwork::V6(self_v6), IpNetwork::V6(other_v6)) => { + IpNetworkIterator::V6(self_v6.sub(other_v6)) + } + (IpNetwork::V4(_), IpNetwork::V6(_)) => { + panic!("Can't remove IPv6 network from IPv4 network") + } + (IpNetwork::V6(_), IpNetwork::V4(_)) => { + panic!("Can't remove IPv4 network from IPv6 network") + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::{ + collections::HashSet, + net::{IpAddr, Ipv4Addr}, + }; + + #[test] + fn subtract_out_of_range() { + let minuend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(25, 0, 0, 0)), 8).unwrap(); + let subtrahend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(125, 92, 4, 0)), 24).unwrap(); + + let difference: Vec<_> = minuend.sub(subtrahend).collect(); + + let expected = vec![minuend]; + + assert_eq!(difference, expected); + } + + #[test] + fn subtract_whole_range() { + let minuend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(25, 0, 0, 0)), 8).unwrap(); + let subtrahend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(16, 0, 0, 0)), 4).unwrap(); + + let difference: Vec<_> = minuend.sub(subtrahend).collect(); + + assert!(difference.is_empty()); + } + + #[test] + fn subtract_inner_range() { + let minuend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 0)), 8).unwrap(); + let subtrahend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(10, 10, 10, 0)), 24).unwrap(); + + let difference: HashSet<_> = minuend.sub(subtrahend).collect(); + + let expected = vec![ + ([10, 0, 0, 0], 13), + ([10, 8, 0, 0], 15), + ([10, 10, 0, 0], 21), + ([10, 10, 8, 0], 23), + ([10, 10, 11, 0], 24), + ([10, 10, 12, 0], 22), + ([10, 10, 16, 0], 20), + ([10, 10, 32, 0], 19), + ([10, 10, 64, 0], 18), + ([10, 10, 128, 0], 17), + ([10, 11, 0, 0], 16), + ([10, 12, 0, 0], 14), + ([10, 16, 0, 0], 12), + ([10, 32, 0, 0], 11), + ([10, 64, 0, 0], 10), + ([10, 128, 0, 0], 9), + ]; + + let expected: HashSet<_> = expected + .into_iter() + .map(|(octets, prefix)| IpNetwork::new(IpAddr::V4(octets.into()), prefix).unwrap()) + .collect(); + + assert_eq!(difference, expected); + } + + #[test] + fn subtract_single_address() { + let minuend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(10, 64, 0, 0)), 10).unwrap(); + let subtrahend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(10, 64, 0, 0)), 32).unwrap(); + + let difference: HashSet<_> = minuend.sub(subtrahend).collect(); + + let expected = vec![ + ([10, 64, 0, 1], 32), + ([10, 64, 0, 2], 31), + ([10, 64, 0, 4], 30), + ([10, 64, 0, 8], 29), + ([10, 64, 0, 16], 28), + ([10, 64, 0, 32], 27), + ([10, 64, 0, 64], 26), + ([10, 64, 0, 128], 25), + ([10, 64, 1, 0], 24), + ([10, 64, 2, 0], 23), + ([10, 64, 4, 0], 22), + ([10, 64, 8, 0], 21), + ([10, 64, 16, 0], 20), + ([10, 64, 32, 0], 19), + ([10, 64, 64, 0], 18), + ([10, 64, 128, 0], 17), + ([10, 65, 0, 0], 16), + ([10, 66, 0, 0], 15), + ([10, 68, 0, 0], 14), + ([10, 72, 0, 0], 13), + ([10, 80, 0, 0], 12), + ([10, 96, 0, 0], 11), + ]; + + let expected: HashSet<_> = expected + .into_iter() + .map(|(octets, prefix)| IpNetwork::new(IpAddr::V4(octets.into()), prefix).unwrap()) + .collect(); + + assert_eq!(difference, expected); + } + + #[test] + fn subtract_multiple() { + let minuend = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), 0).unwrap(); + let subtrahend_1 = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 0)), 8).unwrap(); + let subtrahend_2 = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(172, 16, 0, 0)), 12).unwrap(); + let subtrahend_3 = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(192, 168, 0, 0)), 16).unwrap(); + let subtrahend_4 = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(169, 254, 0, 0)), 16).unwrap(); + let subtrahend_5 = IpNetwork::new(IpAddr::V4(Ipv4Addr::new(224, 0, 0, 0)), 24).unwrap(); + let subtrahend_6 = + IpNetwork::new(IpAddr::V4(Ipv4Addr::new(239, 255, 255, 250)), 32).unwrap(); + let subtrahend_7 = + IpNetwork::new(IpAddr::V4(Ipv4Addr::new(239, 255, 255, 251)), 32).unwrap(); + + let difference: HashSet<_> = minuend + .sub_all(vec![ + subtrahend_1, + subtrahend_2, + subtrahend_3, + subtrahend_4, + subtrahend_5, + subtrahend_6, + subtrahend_7, + ]) + .collect(); + + let expected = vec![ + ([0, 0, 0, 0], 5), + ([8, 0, 0, 0], 7), + ([11, 0, 0, 0], 8), + ([12, 0, 0, 0], 6), + ([16, 0, 0, 0], 4), + ([32, 0, 0, 0], 3), + ([64, 0, 0, 0], 2), + ([128, 0, 0, 0], 3), + ([160, 0, 0, 0], 5), + ([168, 0, 0, 0], 8), + ([169, 0, 0, 0], 9), + ([169, 128, 0, 0], 10), + ([169, 192, 0, 0], 11), + ([169, 224, 0, 0], 12), + ([169, 240, 0, 0], 13), + ([169, 248, 0, 0], 14), + ([169, 252, 0, 0], 15), + ([169, 255, 0, 0], 16), + ([170, 0, 0, 0], 7), + ([172, 0, 0, 0], 12), + ([172, 32, 0, 0], 11), + ([172, 64, 0, 0], 10), + ([172, 128, 0, 0], 9), + ([173, 0, 0, 0], 8), + ([174, 0, 0, 0], 7), + ([176, 0, 0, 0], 4), + ([192, 0, 0, 0], 9), + ([192, 128, 0, 0], 11), + ([192, 160, 0, 0], 13), + ([192, 169, 0, 0], 16), + ([192, 170, 0, 0], 15), + ([192, 172, 0, 0], 14), + ([192, 176, 0, 0], 12), + ([192, 192, 0, 0], 10), + ([193, 0, 0, 0], 8), + ([194, 0, 0, 0], 7), + ([196, 0, 0, 0], 6), + ([200, 0, 0, 0], 5), + ([208, 0, 0, 0], 4), + ([224, 0, 1, 0], 24), + ([224, 0, 2, 0], 23), + ([224, 0, 4, 0], 22), + ([224, 0, 8, 0], 21), + ([224, 0, 16, 0], 20), + ([224, 0, 32, 0], 19), + ([224, 0, 64, 0], 18), + ([224, 0, 128, 0], 17), + ([224, 1, 0, 0], 16), + ([224, 2, 0, 0], 15), + ([224, 4, 0, 0], 14), + ([224, 8, 0, 0], 13), + ([224, 16, 0, 0], 12), + ([224, 32, 0, 0], 11), + ([224, 64, 0, 0], 10), + ([224, 128, 0, 0], 9), + ([225, 0, 0, 0], 8), + ([226, 0, 0, 0], 7), + ([228, 0, 0, 0], 6), + ([232, 0, 0, 0], 6), + ([236, 0, 0, 0], 7), + ([238, 0, 0, 0], 8), + ([239, 0, 0, 0], 9), + ([239, 128, 0, 0], 10), + ([239, 192, 0, 0], 11), + ([239, 224, 0, 0], 12), + ([239, 240, 0, 0], 13), + ([239, 248, 0, 0], 14), + ([239, 252, 0, 0], 15), + ([239, 254, 0, 0], 16), + ([239, 255, 0, 0], 17), + ([239, 255, 128, 0], 18), + ([239, 255, 192, 0], 19), + ([239, 255, 224, 0], 20), + ([239, 255, 240, 0], 21), + ([239, 255, 248, 0], 22), + ([239, 255, 252, 0], 23), + ([239, 255, 254, 0], 24), + ([239, 255, 255, 0], 25), + ([239, 255, 255, 128], 26), + ([239, 255, 255, 192], 27), + ([239, 255, 255, 224], 28), + ([239, 255, 255, 240], 29), + ([239, 255, 255, 248], 31), + ([239, 255, 255, 252], 30), + ([240, 0, 0, 0], 4), + ]; + + let expected: HashSet<_> = expected + .into_iter() + .map(|(octets, prefix)| IpNetwork::new(IpAddr::V4(octets.into()), prefix).unwrap()) + .collect(); + + assert_eq!(difference, expected); + } +} diff --git a/talpid-core/src/tunnel/tun_provider/android.rs b/talpid-core/src/tunnel/tun_provider/android/mod.rs index 487598332c..6de6217a31 100644 --- a/talpid-core/src/tunnel/tun_provider/android.rs +++ b/talpid-core/src/tunnel/tun_provider/android/mod.rs @@ -1,3 +1,6 @@ +mod ipnetwork_sub; + +use self::ipnetwork_sub::IpNetworkSub; use super::TunConfig; use ipnetwork::IpNetwork; use jnix::{ @@ -66,11 +69,12 @@ pub struct AndroidTunProvider { object: GlobalRef, active_tun: Option<File>, last_tun_config: TunConfig, + allow_lan: bool, } impl AndroidTunProvider { /// Create a new AndroidTunProvider interfacing with Android's VpnService. - pub fn new(context: AndroidContext) -> Self { + pub fn new(context: AndroidContext, allow_lan: bool) -> Self { // Initial configuration simply intercepts all packets. The only field that matters is // `routes`, because it determines what must enter the tunnel. All other fields contain // stub values. @@ -83,6 +87,7 @@ impl AndroidTunProvider { IpNetwork::new(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0)), 0) .expect("Invalid IP network prefix for IPv6 address"), ], + required_routes: vec![], mtu: 1380, }; @@ -100,7 +105,20 @@ impl AndroidTunProvider { object: context.vpn_service, active_tun: None, last_tun_config: initial_tun_config, + allow_lan, + } + } + + pub fn set_allow_lan(&mut self, allow_lan: bool) -> Result<(), Error> { + if self.allow_lan != allow_lan { + self.allow_lan = allow_lan; + + if self.active_tun.is_some() { + self.create_tun()?; + } } + + Ok(()) } /// Retrieve a tunnel device with the provided configuration. @@ -246,7 +264,52 @@ impl AndroidTunProvider { .as_raw_fd()) } + fn prepare_tun_config(&self, config: TunConfig) -> TunConfig { + if self.allow_lan { + let (required_ipv4_routes, required_ipv6_routes) = config + .required_routes + .iter() + .cloned() + .partition::<Vec<_>, _>(|route| route.is_ipv4()); + + let (original_lan_ipv4_networks, original_lan_ipv6_networks) = + crate::firewall::ALLOWED_LAN_NETS + .iter() + .chain(crate::firewall::ALLOWED_LAN_MULTICAST_NETS.iter()) + .cloned() + .partition::<Vec<_>, _>(|network| network.is_ipv4()); + + let lan_ipv4_networks = original_lan_ipv4_networks + .into_iter() + .flat_map(|network| network.sub_all(required_ipv4_routes.iter().cloned())) + .collect::<Vec<_>>(); + + let lan_ipv6_networks = original_lan_ipv6_networks + .into_iter() + .flat_map(|network| network.sub_all(required_ipv6_routes.iter().cloned())) + .collect::<Vec<_>>(); + + let routes = config + .routes + .iter() + .flat_map(|&route| { + if route.is_ipv4() { + route.sub_all(lan_ipv4_networks.iter().cloned()) + } else { + route.sub_all(lan_ipv6_networks.iter().cloned()) + } + }) + .collect(); + + TunConfig { routes, ..config } + } else { + config + } + } + fn open_tun(&mut self, config: TunConfig) -> Result<(), Error> { + let actual_config = self.prepare_tun_config(config.clone()); + let env = JnixEnv::from( self.jvm .attach_current_thread_as_daemon() @@ -260,7 +323,7 @@ impl AndroidTunProvider { ) .map_err(|cause| Error::FindMethod("createTun", cause))?; - let java_config = config.clone().into_java(&env); + let java_config = actual_config.clone().into_java(&env); let result = env .call_method_unchecked( self.object.as_obj(), diff --git a/talpid-core/src/tunnel/tun_provider/mod.rs b/talpid-core/src/tunnel/tun_provider/mod.rs index b4751e019f..9ac1e14895 100644 --- a/talpid-core/src/tunnel/tun_provider/mod.rs +++ b/talpid-core/src/tunnel/tun_provider/mod.rs @@ -6,7 +6,7 @@ use std::net::IpAddr; cfg_if! { if #[cfg(target_os = "android")] { - #[path = "android.rs"] + #[path = "android/mod.rs"] mod imp; use self::imp::{AndroidTunProvider, VpnServiceTun}; pub use self::imp::Error; @@ -51,6 +51,11 @@ pub struct TunConfig { )] pub routes: Vec<IpNetwork>, + /// Routes that are required to be configured for the tunnel. + #[cfg(target_os = "android")] + #[jnix(skip)] + pub required_routes: Vec<IpNetwork>, + /// Maximum Transmission Unit in the tunnel. #[cfg_attr(target_os = "android", jnix(map = "|mtu| mtu as i32"))] pub mtu: u16, diff --git a/talpid-core/src/tunnel/wireguard/wireguard_go.rs b/talpid-core/src/tunnel/wireguard/wireguard_go.rs index c569c8727e..f0c79b595d 100644 --- a/talpid-core/src/tunnel/wireguard/wireguard_go.rs +++ b/talpid-core/src/tunnel/wireguard/wireguard_go.rs @@ -234,11 +234,25 @@ impl WgGoTunnel { addresses: config.tunnel.addresses.clone(), dns_servers, routes: routes.collect(), + #[cfg(target_os = "android")] + required_routes: Self::create_required_routes(config), mtu: config.mtu, } } #[cfg(target_os = "android")] + fn create_required_routes(config: &Config) -> Vec<IpNetwork> { + let mut required_routes = vec![IpNetwork::new(IpAddr::V4(config.ipv4_gateway), 32) + .expect("Invalid IPv4 network prefix")]; + + required_routes.extend(config.ipv6_gateway.map(|address| { + IpNetwork::new(IpAddr::V6(address), 128).expect("Invalid IPv6 network prefix") + })); + + required_routes + } + + #[cfg(target_os = "android")] fn bypass_tunnel_sockets( tunnel_device: &mut Tun, handle: i32, diff --git a/talpid-core/src/tunnel_state_machine/connected_state.rs b/talpid-core/src/tunnel_state_machine/connected_state.rs index 3ce80db030..1c0407a615 100644 --- a/talpid-core/src/tunnel_state_machine/connected_state.rs +++ b/talpid-core/src/tunnel_state_machine/connected_state.rs @@ -110,21 +110,23 @@ impl ConnectedState { match try_handle_event!(self, commands.poll()) { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; - - match self.set_firewall_policy(shared_values) { - Ok(()) => SameState(self), - Err(error) => { - log::error!( - "{}", - error.display_chain_with_msg( - "Failed to apply firewall policy for connected state" + if let Err(error_cause) = shared_values.set_allow_lan(allow_lan) { + self.disconnect(shared_values, AfterDisconnect::Block(error_cause)) + } else { + match self.set_firewall_policy(shared_values) { + Ok(()) => SameState(self), + Err(error) => { + log::error!( + "{}", + error.display_chain_with_msg( + "Failed to apply firewall policy for connected state" + ) + ); + self.disconnect( + shared_values, + AfterDisconnect::Block(ErrorStateCause::SetFirewallPolicyError), ) - ); - self.disconnect( - shared_values, - AfterDisconnect::Block(ErrorStateCause::SetFirewallPolicyError), - ) + } } } } diff --git a/talpid-core/src/tunnel_state_machine/connecting_state.rs b/talpid-core/src/tunnel_state_machine/connecting_state.rs index 720e2d3729..848b2c7485 100644 --- a/talpid-core/src/tunnel_state_machine/connecting_state.rs +++ b/talpid-core/src/tunnel_state_machine/connecting_state.rs @@ -162,6 +162,17 @@ impl ConnectingState { } } + fn disconnect( + self, + shared_values: &mut SharedTunnelStateValues, + after_disconnect: AfterDisconnect, + ) -> EventConsequence<Self> { + EventConsequence::NewState(DisconnectingState::enter( + shared_values, + (self.close_handle, self.tunnel_close_event, after_disconnect), + )) + } + fn handle_commands( self, commands: &mut mpsc::UnboundedReceiver<TunnelCommand>, @@ -171,25 +182,24 @@ impl ConnectingState { match try_handle_event!(self, commands.poll()) { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; - match Self::set_firewall_policy(shared_values, &self.tunnel_parameters) { - Ok(()) => SameState(self), - Err(error) => { - error!( - "{}", - error.display_chain_with_msg( - "Failed to apply firewall policy for connecting state" - ) - ); + if let Err(error_cause) = shared_values.set_allow_lan(allow_lan) { + self.disconnect(shared_values, AfterDisconnect::Block(error_cause)) + } else { + match Self::set_firewall_policy(shared_values, &self.tunnel_parameters) { + Ok(()) => SameState(self), + Err(error) => { + error!( + "{}", + error.display_chain_with_msg( + "Failed to apply firewall policy for connecting state" + ) + ); - NewState(DisconnectingState::enter( - shared_values, - ( - self.close_handle, - self.tunnel_close_event, + self.disconnect( + shared_values, AfterDisconnect::Block(ErrorStateCause::SetFirewallPolicyError), - ), - )) + ) + } } } } @@ -200,42 +210,23 @@ impl ConnectingState { Ok(TunnelCommand::IsOffline(is_offline)) => { shared_values.is_offline = is_offline; if is_offline { - NewState(DisconnectingState::enter( + self.disconnect( shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Block(ErrorStateCause::IsOffline), - ), - )) + AfterDisconnect::Block(ErrorStateCause::IsOffline), + ) } else { SameState(self) } } - Ok(TunnelCommand::Connect) => NewState(DisconnectingState::enter( - shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Reconnect(0), - ), - )), - Ok(TunnelCommand::Disconnect) | Err(_) => NewState(DisconnectingState::enter( - shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Nothing, - ), - )), - Ok(TunnelCommand::Block(reason)) => NewState(DisconnectingState::enter( - shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Block(reason), - ), - )), + Ok(TunnelCommand::Connect) => { + self.disconnect(shared_values, AfterDisconnect::Reconnect(0)) + } + Ok(TunnelCommand::Disconnect) | Err(_) => { + self.disconnect(shared_values, AfterDisconnect::Nothing) + } + Ok(TunnelCommand::Block(reason)) => { + self.disconnect(shared_values, AfterDisconnect::Block(reason)) + } } } @@ -247,14 +238,10 @@ impl ConnectingState { use self::EventConsequence::*; match try_handle_event!(self, self.tunnel_events.poll()) { - Ok(TunnelEvent::AuthFailed(reason)) => NewState(DisconnectingState::enter( + Ok(TunnelEvent::AuthFailed(reason)) => self.disconnect( shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Block(ErrorStateCause::AuthFailed(reason)), - ), - )), + AfterDisconnect::Block(ErrorStateCause::AuthFailed(reason)), + ), Ok(TunnelEvent::Up(metadata)) => NewState(ConnectedState::enter( shared_values, self.into_connected_state_bootstrap(metadata), @@ -262,14 +249,8 @@ impl ConnectingState { Ok(_) => SameState(self), Err(_) => { debug!("The tunnel disconnected unexpectedly"); - NewState(DisconnectingState::enter( - shared_values, - ( - self.close_handle, - self.tunnel_close_event, - AfterDisconnect::Reconnect(self.retry_attempt + 1), - ), - )) + let retry_attempt = self.retry_attempt + 1; + self.disconnect(shared_values, AfterDisconnect::Reconnect(retry_attempt)) } } } diff --git a/talpid-core/src/tunnel_state_machine/disconnected_state.rs b/talpid-core/src/tunnel_state_machine/disconnected_state.rs index f183a7c78c..0eb02252dd 100644 --- a/talpid-core/src/tunnel_state_machine/disconnected_state.rs +++ b/talpid-core/src/tunnel_state_machine/disconnected_state.rs @@ -59,7 +59,12 @@ impl TunnelState for DisconnectedState { match try_handle_event!(self, commands.poll()) { Ok(TunnelCommand::AllowLan(allow_lan)) => { if shared_values.allow_lan != allow_lan { - shared_values.allow_lan = allow_lan; + // The only platform that can fail is Android, but Android doesn't support the + // "block when disconnected" option, so the following call never fails. + shared_values + .set_allow_lan(allow_lan) + .expect("Failed to set allow LAN parameter"); + Self::set_firewall_policy(shared_values); } SameState(self) diff --git a/talpid-core/src/tunnel_state_machine/disconnecting_state.rs b/talpid-core/src/tunnel_state_machine/disconnecting_state.rs index d18fdeb6ee..c07ecbf8f7 100644 --- a/talpid-core/src/tunnel_state_machine/disconnecting_state.rs +++ b/talpid-core/src/tunnel_state_machine/disconnecting_state.rs @@ -32,7 +32,7 @@ impl DisconnectingState { self.after_disconnect = match after_disconnect { AfterDisconnect::Nothing => match event { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; + let _ = shared_values.set_allow_lan(allow_lan); AfterDisconnect::Nothing } Ok(TunnelCommand::BlockWhenDisconnected(block_when_disconnected)) => { @@ -49,7 +49,7 @@ impl DisconnectingState { }, AfterDisconnect::Block(reason) => match event { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; + let _ = shared_values.set_allow_lan(allow_lan); AfterDisconnect::Block(reason) } Ok(TunnelCommand::BlockWhenDisconnected(block_when_disconnected)) => { @@ -71,7 +71,7 @@ impl DisconnectingState { }, AfterDisconnect::Reconnect(retry_attempt) => match event { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; + let _ = shared_values.set_allow_lan(allow_lan); AfterDisconnect::Reconnect(retry_attempt) } Ok(TunnelCommand::BlockWhenDisconnected(block_when_disconnected)) => { diff --git a/talpid-core/src/tunnel_state_machine/error_state.rs b/talpid-core/src/tunnel_state_machine/error_state.rs index 9d8402997d..b1bf5183b2 100644 --- a/talpid-core/src/tunnel_state_machine/error_state.rs +++ b/talpid-core/src/tunnel_state_machine/error_state.rs @@ -81,9 +81,12 @@ impl TunnelState for ErrorState { match try_handle_event!(self, commands.poll()) { Ok(TunnelCommand::AllowLan(allow_lan)) => { - shared_values.allow_lan = allow_lan; - Self::set_firewall_policy(shared_values); - SameState(self) + if let Err(error_state_cause) = shared_values.set_allow_lan(allow_lan) { + NewState(Self::enter(shared_values, error_state_cause)) + } else { + Self::set_firewall_policy(shared_values); + SameState(self) + } } Ok(TunnelCommand::BlockWhenDisconnected(block_when_disconnected)) => { shared_values.block_when_disconnected = block_when_disconnected; diff --git a/talpid-core/src/tunnel_state_machine/mod.rs b/talpid-core/src/tunnel_state_machine/mod.rs index 76bd62a1d4..3fb2e4e757 100644 --- a/talpid-core/src/tunnel_state_machine/mod.rs +++ b/talpid-core/src/tunnel_state_machine/mod.rs @@ -89,6 +89,8 @@ where let tun_provider = TunProvider::new( #[cfg(target_os = "android")] android_context, + #[cfg(target_os = "android")] + allow_lan, ); let (startup_result_tx, startup_result_rx) = sync_mpsc::channel(); @@ -324,6 +326,30 @@ struct SharedTunnelStateValues { resource_dir: PathBuf, } +impl SharedTunnelStateValues { + pub fn set_allow_lan(&mut self, allow_lan: bool) -> Result<(), ErrorStateCause> { + if self.allow_lan != allow_lan { + self.allow_lan = allow_lan; + + #[cfg(target_os = "android")] + { + if let Err(error) = self.tun_provider.set_allow_lan(allow_lan) { + log::error!( + "{}", + error.display_chain_with_msg(&format!( + "Failed to restart tunnel after {} LAN connections", + if allow_lan { "allowing" } else { "blocking" } + )) + ); + return Err(ErrorStateCause::StartTunnelError); + } + } + } + + Ok(()) + } +} + /// Asynchronous result of an attempt to progress a state. enum EventConsequence<T: TunnelState> { /// Transition to a new state. |
