use crate::{CredentialRetrieverCreator, Identifier, SecureChannelOptions, TrustPolicy};
use crate::{SecureChannel, SecureChannels};
use ockam_core::api::Reply::Successful;
use ockam_core::api::{Error, Reply, Request, Response};
use ockam_core::compat::sync::Arc;
use ockam_core::compat::time::Duration;
use ockam_core::compat::vec::Vec;
use ockam_core::{self, route, Address, Message, Result, Route};
use ockam_node::api::Client;
use ockam_node::Context;
use ockam_transport_core::Transport;
#[derive(Clone)]
pub struct SecureClient {
secure_channels: Arc<SecureChannels>,
credential_retriever_creator: Option<Arc<dyn CredentialRetrieverCreator>>,
transport: Arc<dyn Transport>,
secure_route: Route,
server_trust_policy: Arc<dyn TrustPolicy>,
client_identifier: Identifier,
secure_channel_timeout: Duration,
request_timeout: Duration,
}
impl SecureClient {
#[allow(clippy::too_many_arguments)]
pub fn new(
secure_channels: Arc<SecureChannels>,
credential_retriever_creator: Option<Arc<dyn CredentialRetrieverCreator>>,
transport: Arc<dyn Transport>,
server_route: Route,
server_trust_policy: Arc<dyn TrustPolicy>,
client_identifier: &Identifier,
secure_channel_timeout: Duration,
request_timeout: Duration,
) -> SecureClient {
Self {
secure_channels,
credential_retriever_creator,
transport,
secure_route: server_route,
server_trust_policy: server_trust_policy.clone(),
client_identifier: client_identifier.clone(),
secure_channel_timeout,
request_timeout,
}
}
pub fn secure_channels(&self) -> Arc<SecureChannels> {
self.secure_channels.clone()
}
pub fn credential_retriever_creator(&self) -> Option<Arc<dyn CredentialRetrieverCreator>> {
self.credential_retriever_creator.clone()
}
pub fn transport(&self) -> Arc<dyn Transport> {
self.transport.clone()
}
pub fn secure_route(&self) -> &Route {
&self.secure_route
}
pub fn server_trust_policy(&self) -> Arc<dyn TrustPolicy> {
self.server_trust_policy.clone()
}
pub fn client_identifier(&self) -> &Identifier {
&self.client_identifier
}
pub fn secure_channel_timeout(&self) -> Duration {
self.secure_channel_timeout
}
pub fn request_timeout(&self) -> Duration {
self.request_timeout
}
pub fn with_secure_channel_timeout(self, timeout: &Duration) -> Self {
Self {
secure_channel_timeout: *timeout,
..self
}
}
pub fn with_request_timeout(self, timeout: &Duration) -> Self {
Self {
request_timeout: *timeout,
..self
}
}
pub fn with_client_identifier(self, client_identifier: &Identifier) -> Self {
Self {
client_identifier: client_identifier.clone(),
..self
}
}
}
impl SecureClient {
pub async fn ask<T, R>(
&self,
ctx: &Context,
api_service: &str,
req: Request<T>,
) -> Result<Reply<R>>
where
T: Message,
R: Message,
{
let response: Response<Vec<u8>> = self
.request_with_timeout(ctx, api_service, req, self.request_timeout)
.await?;
response.to_reply()
}
pub async fn tell<T>(
&self,
ctx: &Context,
api_service: &str,
req: Request<T>,
) -> Result<Reply<()>>
where
T: Message,
{
let request_header = req.header().clone();
let response: Response<Vec<u8>> = self
.request_with_timeout(ctx, api_service, req, self.request_timeout)
.await?;
if response.is_ok() {
Ok(Successful(()))
} else {
let status = response.header().status();
Ok(Reply::Failed(
response
.get_error()
.unwrap_or(Error::from_failed_request(&request_header, "missing error")),
status,
))
}
}
pub async fn request<T, R>(
&self,
ctx: &Context,
api_service: &str,
req: Request<T>,
) -> Result<Response<R>>
where
T: Message,
R: Message,
{
self.request_with_timeout(ctx, api_service, req, self.request_timeout)
.await
}
pub async fn request_with_timeout<T, R>(
&self,
ctx: &Context,
api_service: &str,
req: Request<T>,
timeout: Duration,
) -> Result<Response<R>>
where
T: Message,
R: Message,
{
let (secure_channel, transport_address) = self.create_secure_channel(ctx).await?;
let route = route![secure_channel.clone(), api_service];
let client = Client::new(&route, Some(timeout));
let response = client.request(ctx, req).await;
let _ = self
.secure_channels
.stop_secure_channel(ctx, secure_channel.encryptor_address());
if let Some(transport_address) = transport_address {
let _ = self.transport.disconnect(&transport_address);
}
response
}
pub async fn create_secure_channel(
&self,
ctx: &Context,
) -> Result<(SecureChannel, Option<Address>)> {
let transport_type = self.transport.transport_type();
let (resolved_route, transport_address) = Context::resolve_transport_route_static(
self.secure_route.clone(),
[(transport_type, self.transport.clone())].into(),
)
.await?;
let options = SecureChannelOptions::new()
.with_trust_policy(self.server_trust_policy())
.with_timeout(self.secure_channel_timeout);
let options =
if let Some(credential_retriever_creator) = self.credential_retriever_creator.clone() {
options.with_credential_retriever_creator(credential_retriever_creator)?
} else {
options
};
let secure_channel = self
.secure_channels
.create_secure_channel(ctx, &self.client_identifier, resolved_route, options)
.await?;
Ok((secure_channel, transport_address))
}
pub async fn check_secure_channel(&self, ctx: &Context) -> Result<()> {
let (secure_channel, transport_address) = self.create_secure_channel(ctx).await?;
let _ = self
.secure_channels
.stop_secure_channel(ctx, secure_channel.encryptor_address());
if let Some(transport_address) = transport_address {
let _ = self.transport.disconnect(&transport_address);
}
Ok(())
}
}