diff options
| author | David Lönnhager <david.l@mullvad.net> | 2024-02-15 19:47:07 +0100 |
|---|---|---|
| committer | David Lönnhager <david.l@mullvad.net> | 2024-02-16 16:37:37 +0100 |
| commit | e471d0739446279b01022090ac4457fe337ca598 (patch) | |
| tree | a67ac87a161cae956cde7bc4023cfd12beba5a6a /mullvad-api | |
| parent | c8a3a3be92098cf64bc9269b3c4791e41c3b500d (diff) | |
| download | mullvadvpn-e471d0739446279b01022090ac4457fe337ca598.tar.xz mullvadvpn-e471d0739446279b01022090ac4457fe337ca598.zip | |
Refactor API access methods
Diffstat (limited to 'mullvad-api')
| -rw-r--r-- | mullvad-api/src/bin/relay_list.rs | 10 | ||||
| -rw-r--r-- | mullvad-api/src/lib.rs | 25 | ||||
| -rw-r--r-- | mullvad-api/src/proxy.rs | 42 | ||||
| -rw-r--r-- | mullvad-api/src/rest.rs | 65 |
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. |
