use async_trait::async_trait;
use reqwest::{Request, Response};
use reqwest_middleware::{Middleware, Next};
use crate::client::accept_payment_policy::AcceptPaymentPolicy;
use crate::client::events::{
ChallengeReceivedContext, ClientEventSubscription, ClientEvents, CredentialCreatedContext,
PaymentFailedContext, PaymentResponseContext,
};
use crate::client::flow::{Exchange, FlowError, PaymentFlow};
use crate::client::provider::PaymentProvider;
use crate::client::DEFAULT_MAX_PAYMENT_RETRIES;
pub struct PaymentMiddleware<P> {
provider: P,
accept_payment_policy: AcceptPaymentPolicy,
events: ClientEvents,
max_payment_retries: usize,
}
impl<P> PaymentMiddleware<P> {
pub fn new(provider: P) -> Self {
Self {
provider,
accept_payment_policy: AcceptPaymentPolicy::default(),
events: ClientEvents::default(),
max_payment_retries: DEFAULT_MAX_PAYMENT_RETRIES,
}
}
pub fn with_accept_payment_policy(mut self, policy: AcceptPaymentPolicy) -> Self {
self.accept_payment_policy = policy;
self
}
pub fn with_events(mut self, events: ClientEvents) -> Self {
self.events = events;
self
}
pub fn with_max_payment_retries(mut self, max_payment_retries: usize) -> Self {
self.max_payment_retries = max_payment_retries;
self
}
pub fn events(&self) -> ClientEvents {
self.events.clone()
}
pub fn on_challenge_received<F, Fut>(&self, handler: F) -> ClientEventSubscription
where
F: Fn(ChallengeReceivedContext) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Option<crate::protocol::core::PaymentCredential>>
+ Send
+ 'static,
{
self.events.on_challenge_received(handler)
}
pub fn on_credential_created<F, Fut>(&self, handler: F) -> ClientEventSubscription
where
F: Fn(CredentialCreatedContext) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
self.events.on_credential_created(handler)
}
pub fn on_payment_response<F, Fut>(&self, handler: F) -> ClientEventSubscription
where
F: Fn(PaymentResponseContext) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
self.events.on_payment_response(handler)
}
pub fn on_payment_failed<F, Fut>(&self, handler: F) -> ClientEventSubscription
where
F: Fn(PaymentFailedContext) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
self.events.on_payment_failed(handler)
}
}
#[async_trait]
impl<P> Middleware for PaymentMiddleware<P>
where
P: PaymentProvider + 'static,
{
async fn handle(
&self,
req: Request,
extensions: &mut http_types::Extensions,
next: Next<'_>,
) -> reqwest_middleware::Result<Response> {
let flow = PaymentFlow {
provider: &self.provider,
policy: &self.accept_payment_policy,
events: &self.events,
max_payment_retries: self.max_payment_retries,
};
flow.run(&mut NextExchange { next, extensions }, req, None)
.await
.map_err(|err| match err {
FlowError::Payment(err) => {
reqwest_middleware::Error::Middleware(anyhow::Error::new(err))
}
FlowError::Send(err) => err,
})
}
}
struct NextExchange<'a, 'b> {
next: Next<'a>,
extensions: &'b mut http_types::Extensions,
}
impl Exchange for NextExchange<'_, '_> {
type Error = reqwest_middleware::Error;
async fn send(&mut self, request: Request) -> reqwest_middleware::Result<Response> {
self.next.clone().run(request, self.extensions).await
}
}
#[cfg(test)]
mod tests;