use futures::future::BoxFuture;
use reqwest::header::{HeaderMap, HeaderValue};
use reqwest::Response;
use serde::de::DeserializeOwned;
use std::collections::HashMap;
use std::convert::{TryFrom, TryInto};
use std::sync::Arc;
use tokio::sync::Mutex;
use typed_builder::TypedBuilder;
use url::Url;
use crate::client::states::*;
use crate::error::WWSVCError;
use crate::params::Parameters;
use crate::requests::{ExecJsonRequest, RequestToHttpString, ToServiceFunctionParameters};
use crate::responses::RegisterResponse;
use crate::{AppHash, Credentials, Cursor, WWClientResult};
#[derive(TypedBuilder)]
#[builder(build_method(into = WebwareClient::<Unregistered>))]
pub struct InternalWebwareClient {
#[builder(setter(transform = |url: &str| {
Url::parse(url).expect("Failed to parse URL").join("/WWSVC/").expect("Failed to join URL")
}))]
webware_url: Url,
#[builder(setter(transform = |vendor_hash: &str| vendor_hash.to_string()))]
vendor_hash: String,
#[builder(setter(transform = |app_hash: &str| app_hash.to_string()))]
app_hash: String,
#[builder(setter(transform = |app_secret: &str| app_secret.to_string()))]
secret: String,
revision: u32,
#[builder(default, setter(transform = |credentials: Credentials| Some(credentials)))]
credentials: Option<Credentials>,
#[builder(default = 1000)]
result_max_lines: u32,
#[builder(default = false)]
allow_insecure: bool,
#[builder(default = std::time::Duration::from_secs(60))]
timeout: std::time::Duration,
}
pub mod states {
#[derive(Clone)]
pub struct Unregistered;
#[derive(Clone)]
pub struct Registered;
pub trait Ready: Send + Sync {}
impl Ready for Registered {}
}
#[derive(Debug)]
struct MutableClientState {
result_max_lines: u32,
cursor: Option<Cursor>,
current_request: u32,
suspend_cursor: bool,
}
#[derive(Clone)]
pub struct WebwareClient<State = Unregistered> {
webware_url: Url,
vendor_hash: String,
app_hash: String,
secret: String,
revision: u32,
credentials: Option<Credentials>,
mutable_state: Arc<Mutex<MutableClientState>>,
client: reqwest::Client,
state: std::marker::PhantomData<State>,
}
impl From<InternalWebwareClient> for WebwareClient<Unregistered> {
fn from(client: InternalWebwareClient) -> Self {
let req_client = reqwest::Client::builder()
.danger_accept_invalid_certs(client.allow_insecure)
.timeout(client.timeout)
.build()
.expect("Failed to build client");
WebwareClient {
webware_url: client.webware_url,
vendor_hash: client.vendor_hash,
app_hash: client.app_hash,
secret: client.secret,
revision: client.revision,
credentials: client.credentials,
mutable_state: Arc::new(Mutex::new(MutableClientState {
result_max_lines: client.result_max_lines,
cursor: None,
current_request: 0,
suspend_cursor: false,
})),
client: req_client,
state: std::marker::PhantomData::<Unregistered>,
}
}
}
impl TryFrom<InternalWebwareClient> for WebwareClient<Registered> {
type Error = WWSVCError;
fn try_from(client: InternalWebwareClient) -> Result<Self, Self::Error> {
let req_client = reqwest::Client::builder()
.danger_accept_invalid_certs(client.allow_insecure)
.timeout(client.timeout)
.build()
.expect("Failed to build client");
if client.credentials.is_none() {
return Err(WWSVCError::MissingCredentials);
}
Ok(WebwareClient {
webware_url: client.webware_url,
vendor_hash: client.vendor_hash,
app_hash: client.app_hash,
secret: client.secret,
revision: client.revision,
credentials: client.credentials,
mutable_state: Arc::new(Mutex::new(MutableClientState {
result_max_lines: client.result_max_lines,
cursor: None,
current_request: 0,
suspend_cursor: false,
})),
client: req_client,
state: std::marker::PhantomData::<Registered>,
})
}
}
impl WebwareClient {
pub fn builder() -> InternalWebwareClientBuilder {
InternalWebwareClient::builder()
}
pub async fn register(self) -> WWClientResult<WebwareClient<Registered>> {
if self.credentials.is_some() {
return Ok(WebwareClient {
webware_url: self.webware_url,
vendor_hash: self.vendor_hash,
app_hash: self.app_hash,
secret: self.secret,
revision: self.revision,
credentials: self.credentials,
mutable_state: self.mutable_state,
client: self.client,
state: std::marker::PhantomData::<Registered>,
});
}
let target_url = self
.webware_url
.join("WWSERVICE/")?
.join("REGISTER/")?
.join(&format!("{}/", self.vendor_hash))?
.join(&format!("{}/", self.app_hash))?
.join(&format!("{}/", self.secret))?
.join(&format!("{}/", self.revision))?;
let response = self.client.get(target_url).send().await?;
let response_obj = response.json::<RegisterResponse>().await?;
Ok(WebwareClient {
webware_url: self.webware_url,
vendor_hash: self.vendor_hash,
app_hash: self.app_hash,
secret: self.secret,
revision: self.revision,
credentials: Some(Credentials {
service_pass: response_obj.service_pass.pass_id,
app_id: response_obj.service_pass.app_id,
}),
mutable_state: self.mutable_state,
client: self.client,
state: std::marker::PhantomData::<Registered>,
})
}
pub async fn with_registered<F, T>(self, f: F) -> WWClientResult<T>
where
F: for<'a> FnOnce(&'a WebwareClient<Registered>) -> BoxFuture<'a, T>,
{
let client = self.register().await?;
let result = f(&client).await;
let _ = client.deregister().await?;
Ok(result)
}
}
impl<State: Ready> WebwareClient<State> {
pub async fn create_cursor(&self, max_lines: u32) {
let cursor = Cursor::new(max_lines);
let mut state = self.mutable_state.lock().await;
state.cursor = Some(cursor);
state.result_max_lines = max_lines;
}
pub async fn close_cursor(&self) {
let mut state = self.mutable_state.lock().await;
state.cursor = None;
}
pub async fn has_cursor(&self) -> bool {
let state = self.mutable_state.lock().await;
state.cursor.is_some()
}
pub fn credentials(&self) -> &Credentials {
self.credentials.as_ref().unwrap()
}
pub async fn set_result_max_lines(&self, max_lines: u32) {
let mut state = self.mutable_state.lock().await;
state.result_max_lines = max_lines;
}
pub async fn get_default_headers(
&self,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<HeaderMap> {
let mut state = self.mutable_state.lock().await;
let mut max_lines = state.result_max_lines;
let mut headers = HashMap::new();
if let Some(credentials) = &self.credentials {
let app_hash = AppHash::new(state.current_request, &credentials.app_id);
state.current_request = app_hash.request_id;
headers.insert("WWSVC-REQID".to_string(), format!("{}", state.current_request));
headers.insert("WWSVC-TS".to_string(), app_hash.date_formatted.to_string());
headers.insert("WWSVC-HASH".to_string(), format!("{:x}", app_hash));
if !state.suspend_cursor {
if let Some(cursor) = &state.cursor {
if !Cursor::closed(cursor) {
headers.insert("WWSVC-CURSOR".to_string(), cursor.cursor_id.to_string());
max_lines = cursor.max_lines;
}
}
}
}
headers.insert("WWSVC-EXECUTE-MODE".to_string(), "SYNCHRON".to_string());
headers.insert("WWSVC-ACCEPT-RESULT-TYPE".to_string(), "JSON".to_string());
headers.insert("WWSVC-ACCEPT-RESULT-MAX-LINES".to_string(), max_lines.to_string());
if let Some(additional_headers) = additional_headers {
headers.extend(
additional_headers
.iter()
.map(|(s1, s2)| (s1.to_string(), s2.to_string())),
);
}
(&headers).try_into().map_err(|_| WWSVCError::InvalidHeader)
}
pub async fn get_bin_headers(
&self,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<HeaderMap> {
let mut headers = self.get_default_headers(additional_headers).await?;
headers.remove("WWSVC-ACCEPT-RESULT-TYPE");
headers.append("WWSVC-ACCEPT-RESULT-TYPE", HeaderValue::from_str("BIN")?);
Ok(headers)
}
pub async fn deregister(self) -> WWClientResult<WebwareClient<Unregistered>> {
if let Some(credentials) = &self.credentials {
let target_url = self
.webware_url
.join("WWSERVICE/")?
.join("DEREGISTER/")?
.join(&format!("{}/", &credentials.service_pass))?;
let headers = self.get_default_headers(None).await?;
let _ = self.client.get(target_url).headers(headers).send().await;
}
Ok(WebwareClient {
webware_url: self.webware_url,
vendor_hash: self.vendor_hash,
app_hash: self.app_hash,
secret: self.secret,
revision: self.revision,
credentials: None,
mutable_state: self.mutable_state,
client: self.client,
state: std::marker::PhantomData::<Unregistered>,
})
}
pub async fn prepare_request(
&self,
method: reqwest::Method,
function: &str,
version: u32,
parameters: Parameters,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<reqwest::Request> {
if self.credentials.is_none() {
return Err(WWSVCError::NotAuthenticated);
}
let target_url = self.webware_url.join("EXECJSON")?;
let headers = self.get_default_headers(additional_headers).await?;
let app_hash_header = headers.get("WWSVC-HASH");
let timestamp_header = headers.get("WWSVC-TS");
let app_hash: String = app_hash_header
.unwrap_or(&HeaderValue::from_str("").unwrap())
.to_str()
.map_err(|_| WWSVCError::HeaderValueToStrError)?
.to_string();
let timestamp: String = timestamp_header
.unwrap_or(&HeaderValue::from_str("").unwrap())
.to_str()
.map_err(|_| WWSVCError::HeaderValueToStrError)?
.to_string();
let parameters = parameters.to_service_function_parameters();
let current_request = {
let state = self.mutable_state.lock().await;
state.current_request
};
let body = ExecJsonRequest::new(
function,
parameters,
version,
&self.credentials.as_ref().unwrap().service_pass,
&app_hash,
×tamp,
current_request,
);
let request = self
.client
.request(method, target_url)
.headers(headers)
.json(&body)
.build()?;
Ok(request)
}
pub async fn execute_request(&self, request: reqwest::Request) -> WWClientResult<Response> {
let response = self.client.execute(request).await?;
let mut state = self.mutable_state.lock().await;
if !state.suspend_cursor {
if let Some(cursor) = &mut state.cursor {
if !Cursor::closed(cursor) && response.headers().contains_key("WWSVC-CURSOR") {
cursor.set_cursor_id(
response
.headers()
.get("WWSVC-CURSOR")
.unwrap()
.to_str()
.unwrap()
.to_string(),
);
}
}
}
Ok(response)
}
pub async fn request(
&self,
method: reqwest::Method,
function: &str,
version: u32,
parameters: Parameters,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<serde_json::Value> {
self.request_generic::<serde_json::Value>(
method,
function,
version,
parameters,
additional_headers,
)
.await
}
pub async fn request_as_response(
&self,
method: reqwest::Method,
function: &str,
version: u32,
parameters: Parameters,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<Response> {
let request =
self.prepare_request(method, function, version, parameters, additional_headers).await?;
tracing::debug!(request = request.to_http_string().unwrap_or_default(), "send request");
let response = self.client.execute(request).await?;
let mut state = self.mutable_state.lock().await;
if !state.suspend_cursor {
if let Some(cursor) = &mut state.cursor {
if !Cursor::closed(cursor) && response.headers().contains_key("WWSVC-CURSOR") {
cursor.set_cursor_id(
response
.headers()
.get("WWSVC-CURSOR")
.unwrap()
.to_str()
.unwrap()
.to_string(),
);
}
}
}
Ok(response)
}
pub async fn request_generic<T>(
&self,
method: reqwest::Method,
function: &str,
version: u32,
parameters: Parameters,
additional_headers: Option<HashMap<&str, &str>>,
) -> WWClientResult<T>
where
T: DeserializeOwned,
{
let response = self
.request_as_response(method, function, version, parameters, additional_headers)
.await?;
let response_obj = response.json::<T>().await?;
Ok(response_obj)
}
pub async fn suspend_cursor(&self) {
let mut state = self.mutable_state.lock().await;
state.suspend_cursor = true;
}
pub async fn resume_cursor(&self) {
let mut state = self.mutable_state.lock().await;
state.suspend_cursor = false;
}
pub async fn cursor_closed(&self) -> bool {
let state = self.mutable_state.lock().await;
state.cursor.as_ref().map_or(true, |c| Cursor::closed(c))
}
}