summaryrefslogtreecommitdiffhomepage
path: root/mullvad-api
diff options
context:
space:
mode:
authorDavid Lönnhager <david.l@mullvad.net>2024-02-19 09:27:24 +0100
committerDavid Lönnhager <david.l@mullvad.net>2024-02-19 09:27:24 +0100
commit10b4df8f352d98e2e3a5507938fbaae4b78d08b6 (patch)
treea67ac87a161cae956cde7bc4023cfd12beba5a6a /mullvad-api
parentc8a3a3be92098cf64bc9269b3c4791e41c3b500d (diff)
parente471d0739446279b01022090ac4457fe337ca598 (diff)
downloadmullvadvpn-10b4df8f352d98e2e3a5507938fbaae4b78d08b6.tar.xz
mullvadvpn-10b4df8f352d98e2e3a5507938fbaae4b78d08b6.zip
Merge branch 'refactor-access-mode-tx' into main
Diffstat (limited to 'mullvad-api')
-rw-r--r--mullvad-api/src/bin/relay_list.rs10
-rw-r--r--mullvad-api/src/lib.rs25
-rw-r--r--mullvad-api/src/proxy.rs42
-rw-r--r--mullvad-api/src/rest.rs65
4 files changed, 82 insertions, 60 deletions
diff --git a/mullvad-api/src/bin/relay_list.rs b/mullvad-api/src/bin/relay_list.rs
index e395d8ae5f..8cb615d77f 100644
--- a/mullvad-api/src/bin/relay_list.rs
+++ b/mullvad-api/src/bin/relay_list.rs
@@ -11,12 +11,10 @@ async fn main() {
let runtime = mullvad_api::Runtime::new(tokio::runtime::Handle::current())
.expect("Failed to load runtime");
- let relay_list_request = RelayListProxy::new(runtime.mullvad_rest_handle(
- ApiConnectionMode::Direct,
- ApiConnectionMode::Direct.into_repeat(),
- ))
- .relay_list(None)
- .await;
+ let relay_list_request =
+ RelayListProxy::new(runtime.mullvad_rest_handle(ApiConnectionMode::Direct.into_provider()))
+ .relay_list(None)
+ .await;
let relay_list = match relay_list_request {
Ok(relay_list) => relay_list,
diff --git a/mullvad-api/src/lib.rs b/mullvad-api/src/lib.rs
index dad6cdf706..6114bec90a 100644
--- a/mullvad-api/src/lib.rs
+++ b/mullvad-api/src/lib.rs
@@ -1,6 +1,5 @@
#[cfg(target_os = "android")]
use futures::channel::mpsc;
-use futures::Stream;
use hyper::Method;
#[cfg(target_os = "android")]
use mullvad_types::account::{PlayPurchase, PlayPurchasePaymentToken};
@@ -8,7 +7,7 @@ use mullvad_types::{
account::{AccountData, AccountToken, VoucherSubmission},
version::AppVersion,
};
-use proxy::ApiConnectionMode;
+use proxy::{ApiConnectionMode, ConnectionModeProvider};
use std::{
cell::Cell,
collections::BTreeMap,
@@ -408,34 +407,30 @@ impl Runtime {
}
/// Creates a new request service and returns a handle to it.
- fn new_request_service<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static>(
+ fn new_request_service<T: ConnectionModeProvider + 'static>(
&self,
sni_hostname: Option<String>,
- initial_connection_mode: ApiConnectionMode,
- proxy_provider: T,
+ connection_mode_provider: T,
#[cfg(target_os = "android")] socket_bypass_tx: Option<mpsc::Sender<SocketBypassRequest>>,
) -> rest::RequestServiceHandle {
rest::RequestService::spawn(
sni_hostname,
self.api_availability.handle(),
self.address_cache.clone(),
- initial_connection_mode,
- proxy_provider,
+ connection_mode_provider,
#[cfg(target_os = "android")]
socket_bypass_tx,
)
}
/// Returns a request factory initialized to create requests for the master API
- pub fn mullvad_rest_handle<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static>(
+ pub fn mullvad_rest_handle<T: ConnectionModeProvider + 'static>(
&self,
- initial_connection_mode: ApiConnectionMode,
- proxy_provider: T,
+ connection_mode_provider: T,
) -> rest::MullvadRestHandle {
let service = self.new_request_service(
Some(API.host().to_string()),
- initial_connection_mode,
- proxy_provider,
+ connection_mode_provider,
#[cfg(target_os = "android")]
self.socket_bypass_tx.clone(),
);
@@ -454,8 +449,7 @@ impl Runtime {
pub fn static_mullvad_rest_handle(&self, hostname: String) -> rest::MullvadRestHandle {
let service = self.new_request_service(
Some(hostname.clone()),
- ApiConnectionMode::Direct,
- futures::stream::repeat(ApiConnectionMode::Direct),
+ ApiConnectionMode::Direct.into_provider(),
#[cfg(target_os = "android")]
self.socket_bypass_tx.clone(),
);
@@ -474,8 +468,7 @@ impl Runtime {
pub fn rest_handle(&self) -> rest::RequestServiceHandle {
self.new_request_service(
None,
- ApiConnectionMode::Direct,
- ApiConnectionMode::Direct.into_repeat(),
+ ApiConnectionMode::Direct.into_provider(),
#[cfg(target_os = "android")]
None,
)
diff --git a/mullvad-api/src/proxy.rs b/mullvad-api/src/proxy.rs
index 2b4821ba64..0915d1d23c 100644
--- a/mullvad-api/src/proxy.rs
+++ b/mullvad-api/src/proxy.rs
@@ -1,4 +1,3 @@
-use futures::Stream;
use hyper::client::connect::Connected;
use serde::{Deserialize, Serialize};
use std::{
@@ -18,6 +17,41 @@ use tokio::{
const CURRENT_CONFIG_FILENAME: &str = "api-endpoint.json";
+pub trait ConnectionModeProvider: Send {
+ /// Initial connection mode
+ fn initial(&self) -> ApiConnectionMode;
+
+ /// Request a new connection mode from the provider
+ fn rotate(&self) -> impl std::future::Future<Output = ()> + Send;
+
+ /// Receive changes to the connection mode, announced by the provider
+ fn receive(&mut self) -> impl std::future::Future<Output = Option<ApiConnectionMode>> + Send;
+}
+
+pub struct StaticConnectionModeProvider {
+ mode: ApiConnectionMode,
+}
+
+impl StaticConnectionModeProvider {
+ pub fn new(mode: ApiConnectionMode) -> Self {
+ Self { mode }
+ }
+}
+
+impl ConnectionModeProvider for StaticConnectionModeProvider {
+ fn initial(&self) -> ApiConnectionMode {
+ self.mode.clone()
+ }
+
+ fn rotate(&self) -> impl std::future::Future<Output = ()> + Send {
+ futures::future::ready(())
+ }
+
+ fn receive(&mut self) -> impl std::future::Future<Output = Option<ApiConnectionMode>> + Send {
+ futures::future::pending()
+ }
+}
+
#[derive(Clone, Debug, PartialEq, Deserialize, Serialize)]
pub enum ApiConnectionMode {
/// Connect directly to the target.
@@ -153,10 +187,8 @@ impl ApiConnectionMode {
*self != ApiConnectionMode::Direct
}
- /// Convenience function that returns a stream that repeats
- /// this config forever.
- pub fn into_repeat(self) -> impl Stream<Item = ApiConnectionMode> {
- futures::stream::repeat(self)
+ pub fn into_provider(self) -> StaticConnectionModeProvider {
+ StaticConnectionModeProvider::new(self)
}
}
diff --git a/mullvad-api/src/rest.rs b/mullvad-api/src/rest.rs
index ca63f16c1f..158d84f01b 100644
--- a/mullvad-api/src/rest.rs
+++ b/mullvad-api/src/rest.rs
@@ -5,12 +5,11 @@ use crate::{
address_cache::AddressCache,
availability::ApiAvailabilityHandle,
https_client_with_sni::{HttpsConnectorWithSni, HttpsConnectorWithSniHandle},
- proxy::ApiConnectionMode,
+ proxy::ConnectionModeProvider,
};
use futures::{
channel::{mpsc, oneshot},
stream::StreamExt,
- Stream,
};
use hyper::{
client::{connect::Connect, Client},
@@ -120,23 +119,22 @@ impl Error {
/// A service that executes HTTP requests, allowing for on-demand termination of all in-flight
/// requests
-pub(crate) struct RequestService<T: Stream<Item = ApiConnectionMode>> {
+pub(crate) struct RequestService<T: ConnectionModeProvider> {
command_tx: Weak<mpsc::UnboundedSender<RequestCommand>>,
command_rx: mpsc::UnboundedReceiver<RequestCommand>,
connector_handle: HttpsConnectorWithSniHandle,
client: hyper::Client<HttpsConnectorWithSni, hyper::Body>,
- proxy_config_provider: T,
+ connection_mode_provider: T,
api_availability: ApiAvailabilityHandle,
}
-impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestService<T> {
+impl<T: ConnectionModeProvider + 'static> RequestService<T> {
/// Constructs a new request service.
pub fn spawn(
sni_hostname: Option<String>,
api_availability: ApiAvailabilityHandle,
address_cache: AddressCache,
- initial_connection_mode: ApiConnectionMode,
- proxy_config_provider: T,
+ connection_mode_provider: T,
#[cfg(target_os = "android")] socket_bypass_tx: Option<mpsc::Sender<SocketBypassRequest>>,
) -> RequestServiceHandle {
let (connector, connector_handle) = HttpsConnectorWithSni::new(
@@ -146,7 +144,7 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
socket_bypass_tx.clone(),
);
- connector_handle.set_connection_mode(initial_connection_mode);
+ connector_handle.set_connection_mode(connection_mode_provider.initial());
let (command_tx, command_rx) = mpsc::unbounded();
let client = Client::builder().build(connector);
@@ -158,7 +156,7 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
command_rx,
connector_handle,
client,
- proxy_config_provider,
+ connection_mode_provider,
api_availability,
};
let handle = RequestServiceHandle { tx: command_tx };
@@ -166,6 +164,27 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
handle
}
+ async fn into_future(mut self) {
+ loop {
+ tokio::select! {
+ new_mode = self.connection_mode_provider.receive() => {
+ let Some(new_mode) = new_mode else {
+ break;
+ };
+ self.connector_handle.set_connection_mode(new_mode);
+ }
+ command = self.command_rx.next() => {
+ let Some(command) = command else {
+ break;
+ };
+
+ self.process_command(command).await;
+ }
+ }
+ }
+ self.connector_handle.reset();
+ }
+
async fn process_command(&mut self, command: RequestCommand) {
match command {
RequestCommand::NewRequest(request, completion_tx) => {
@@ -174,11 +193,8 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
RequestCommand::Reset => {
self.connector_handle.reset();
}
- RequestCommand::NextApiConfig(completion_tx) => {
- if let Some(connection_mode) = self.proxy_config_provider.next().await {
- self.connector_handle.set_connection_mode(connection_mode);
- }
- let _ = completion_tx.send(Ok(()));
+ RequestCommand::NextApiConfig => {
+ self.connection_mode_provider.rotate().await;
}
}
}
@@ -201,8 +217,7 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
if err.is_network_error() && !api_availability.get_state().is_offline() {
log::error!("{}", err.display_chain_with_msg("HTTP request failed"));
if let Some(tx) = tx {
- let (completion_tx, _completion_rx) = oneshot::channel();
- let _ = tx.unbounded_send(RequestCommand::NextApiConfig(completion_tx));
+ let _ = tx.unbounded_send(RequestCommand::NextApiConfig);
}
}
}
@@ -210,13 +225,6 @@ impl<T: Stream<Item = ApiConnectionMode> + Unpin + Send + 'static> RequestServic
let _ = completion_tx.send(response);
});
}
-
- async fn into_future(mut self) {
- while let Some(command) = self.command_rx.next().await {
- self.process_command(command).await;
- }
- self.connector_handle.reset();
- }
}
#[derive(Clone)]
@@ -239,15 +247,6 @@ impl RequestServiceHandle {
.map_err(|_| Error::RestServiceDown)?;
completion_rx.await.map_err(|_| Error::RestServiceDown)?
}
-
- /// Forcibly update the connection mode.
- pub async fn next_api_endpoint(&self) -> Result<()> {
- let (completion_tx, completion_rx) = oneshot::channel();
- self.tx
- .unbounded_send(RequestCommand::NextApiConfig(completion_tx))
- .map_err(|_| Error::RestServiceDown)?;
- completion_rx.await.map_err(|_| Error::RestServiceDown)?
- }
}
#[derive(Debug)]
@@ -257,7 +256,7 @@ pub(crate) enum RequestCommand {
oneshot::Sender<std::result::Result<Response, Error>>,
),
Reset,
- NextApiConfig(oneshot::Sender<std::result::Result<(), Error>>),
+ NextApiConfig,
}
/// A REST request that is sent to the RequestService to be executed.