summaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authorJanito Vaqueiro Ferreira Filho <janito@mullvad.net>2020-02-06 09:28:32 -0300
committerJanito Vaqueiro Ferreira Filho <janito@mullvad.net>2020-02-06 09:28:32 -0300
commit8ced6b2a795352e99d212292c17378a7885b943a (patch)
tree0a89a8d3076d0da44de31ec84e1adfe085aa16c7
parent982b11d0084a205e1305c621ce430bb5f1040b6d (diff)
parent00701395ea0625789ad915c532a4c2417a6985b0 (diff)
downloadmullvadvpn-8ced6b2a795352e99d212292c17378a7885b943a.tar.xz
mullvadvpn-8ced6b2a795352e99d212292c17378a7885b943a.zip
Merge branch 'allow-lan-on-android'
-rw-r--r--talpid-core/src/firewall/mod.rs4
-rw-r--r--talpid-core/src/tunnel/tun_provider/android/ipnetwork_sub.rs452
-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.rs7
-rw-r--r--talpid-core/src/tunnel/wireguard/wireguard_go.rs14
-rw-r--r--talpid-core/src/tunnel_state_machine/connected_state.rs30
-rw-r--r--talpid-core/src/tunnel_state_machine/connecting_state.rs107
-rw-r--r--talpid-core/src/tunnel_state_machine/disconnected_state.rs7
-rw-r--r--talpid-core/src/tunnel_state_machine/disconnecting_state.rs6
-rw-r--r--talpid-core/src/tunnel_state_machine/error_state.rs9
-rw-r--r--talpid-core/src/tunnel_state_machine/mod.rs26
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.