diff options
Diffstat (limited to 'crates/daemon/src/aclient/api_client.rs')
| -rw-r--r-- | crates/daemon/src/aclient/api_client.rs | 197 |
1 files changed, 197 insertions, 0 deletions
diff --git a/crates/daemon/src/aclient/api_client.rs b/crates/daemon/src/aclient/api_client.rs new file mode 100644 index 00000000..666e6dce --- /dev/null +++ b/crates/daemon/src/aclient/api_client.rs @@ -0,0 +1,197 @@ +use std::env; +use std::time::Duration; + +use eyre::{Result, bail, eyre}; +use reqwest::{Response, StatusCode, Url, header::HeaderMap}; +use tracing::debug; +use uuid::Uuid; + +use turtle_common::{api::ErrorResponse, record::RecordStatus}; +use turtle_common::{ + api::{ATUIN_CARGO_VERSION, ATUIN_HEADER_VERSION, ATUIN_VERSION}, + record::{EncryptedData, HostId, Record, RecordIdx}, + tls::ensure_crypto_provider, +}; + +use semver::Version; + +static APP_USER_AGENT: &str = concat!("atuin/", env!("CARGO_PKG_VERSION"),); + +pub(crate) struct Client<'a> { + sync_addr: &'a str, + user_id: Uuid, + inner: reqwest::Client, +} + +fn make_url(address: &str, path: &str, user_id: Uuid) -> Result<String> { + let address = address.strip_suffix('/').unwrap_or(address); + + // `join()` expects a trailing `/` in order to join paths + // e.g. it treats `http://host:port/subdir` as a file called `subdir` + let address = &format!("{address}/api/v0/{user_id}/"); + + // passing a path with a leading `/` will cause `join()` to replace the entire URL path + let path = path.strip_prefix("/").unwrap_or(path); + + let url = Url::parse(address) + .map(|url| url.join(path))? + .map_err(|_| eyre!("invalid address"))?; + + Ok(url.to_string()) +} + +pub(crate) fn ensure_version(response: &Response) -> Result<bool> { + let version = response.headers().get(ATUIN_HEADER_VERSION); + + let version = if let Some(version) = version { + match version.to_str() { + Ok(v) => Version::parse(v), + Err(e) => bail!("failed to parse server version: {:?}", e), + } + } else { + bail!("Server not reporting its version: it is either too old or unhealthy"); + }?; + + // If the client is newer than the server + if version.major < ATUIN_VERSION.major { + println!( + "Atuin version mismatch! In order to successfully sync, the server needs to run a newer version of Atuin" + ); + println!("Client: {ATUIN_CARGO_VERSION}"); + println!("Server: {version}"); + + return Ok(false); + } + + Ok(true) +} + +async fn handle_resp_error(resp: Response) -> Result<Response> { + let status = resp.status(); + let url = resp.url().to_string(); + + if status == StatusCode::SERVICE_UNAVAILABLE { + bail!( + "Service unavailable: check https://status.atuin.sh (or get in touch with your host)" + ); + } + + if status == StatusCode::TOO_MANY_REQUESTS { + bail!("Rate limited; please wait before doing that again"); + } + + if !status.is_success() { + if let Ok(error) = resp.json::<ErrorResponse<'_>>().await { + let reason = error.reason; + + if status.is_client_error() { + bail!("Invalid request to the service at {url}, {status} - {reason}.") + } + + bail!( + "There was an error with the atuin sync service at {url}, server error {status}: {reason}.\nIf the problem persists, contact the host" + ) + } + + bail!( + "There was an error with the atuin sync service at {url}, Status {status:?}.\nIf the problem persists, contact the host" + ) + } + + Ok(resp) +} + +impl<'a> Client<'a> { + pub(crate) fn new( + sync_addr: &'a str, + connect_timeout: u64, + timeout: u64, + user_id: Uuid, + ) -> Result<Self> { + ensure_crypto_provider(); + let mut headers = HeaderMap::new(); + + // used for semver server check + headers.insert(ATUIN_HEADER_VERSION, ATUIN_CARGO_VERSION.parse()?); + + Ok(Client { + user_id, + sync_addr, + inner: reqwest::Client::builder() + .user_agent(APP_USER_AGENT) + .default_headers(headers) + .connect_timeout(Duration::new(connect_timeout, 0)) + .timeout(Duration::new(timeout, 0)) + .build()?, + }) + } + + pub(crate) async fn delete_store(&self) -> Result<()> { + let url = make_url(self.sync_addr, "/store", self.user_id)?; + let url = Url::parse(url.as_str())?; + + let resp = self.inner.delete(url).send().await?; + + handle_resp_error(resp).await?; + + Ok(()) + } + + pub(crate) async fn post_records(&self, records: &[Record<EncryptedData>]) -> Result<()> { + let url = make_url(self.sync_addr, "/record", self.user_id)?; + let url = Url::parse(url.as_str())?; + + debug!("uploading {} records to {url}", records.len()); + + let resp = self.inner.post(url).json(records).send().await?; + handle_resp_error(resp).await?; + + Ok(()) + } + + pub(crate) async fn next_records( + &self, + host: HostId, + tag: String, + start: RecordIdx, + count: u64, + ) -> Result<Vec<Record<EncryptedData>>> { + debug!("fetching record/s from host {}/{}/{}", host.0, tag, start); + + let url = make_url( + self.sync_addr, + &format!( + "/record/next?host={}&tag={}&count={}&start={}", + host.0, tag, count, start + ), + self.user_id, + )?; + + let url = Url::parse(url.as_str())?; + + let resp = self.inner.get(url).send().await?; + let resp = handle_resp_error(resp).await?; + + let records = resp.json::<Vec<Record<EncryptedData>>>().await?; + + Ok(records) + } + + pub(crate) async fn record_status(&self) -> Result<RecordStatus> { + let url = make_url(self.sync_addr, "/record", self.user_id)?; + let url = Url::parse(url.as_str())?; + + let resp = self.inner.get(url).send().await?; + let resp = handle_resp_error(resp).await?; + + if !ensure_version(&resp)? { + bail!("could not sync records due to version mismatch"); + } + + let index = resp.json().await?; + + debug!("got remote index {index:?}"); + + Ok(index) + } +} |
