use std::collections::HashMap;
use std::fmt::{self, Debug, Formatter};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use compact_str::CompactString;
use crate::wire::{
ExtensionEntry, Extensions, PaymentPayload, PaymentRequirements, SettleResponse, VerifyResponse,
};
#[cfg(feature = "ext-bazaar")]
#[cfg_attr(docsrs, doc(cfg(feature = "ext-bazaar")))]
pub mod bazaar;
#[cfg(feature = "ext-payment-id")]
#[cfg_attr(docsrs, doc(cfg(feature = "ext-payment-id")))]
pub mod payment_id;
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[derive(Debug)]
#[non_exhaustive]
pub struct AdvertiseContext<'a> {
pub requirement: Option<&'a PaymentRequirements>,
}
#[non_exhaustive]
pub struct VerifyContext<'a> {
pub payload: &'a PaymentPayload<serde_json::Value, serde_json::Value>,
pub requirements: &'a PaymentRequirements,
pub response: &'a VerifyResponse,
}
impl Debug for VerifyContext<'_> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("VerifyContext")
.field("requirements", &self.requirements)
.finish_non_exhaustive()
}
}
#[non_exhaustive]
pub struct SettleContext<'a> {
pub payload: &'a PaymentPayload<serde_json::Value, serde_json::Value>,
pub requirements: &'a PaymentRequirements,
pub response: &'a SettleResponse,
}
impl Debug for SettleContext<'_> {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("SettleContext")
.field("requirements", &self.requirements)
.finish_non_exhaustive()
}
}
pub trait Extension: Send + Sync {
fn id(&self) -> &'static str;
fn advertise(&self, _ctx: &AdvertiseContext<'_>) -> Option<ExtensionEntry> {
None
}
fn on_verify<'a>(
&'a self,
_ctx: &'a VerifyContext<'a>,
) -> impl Future<Output = Option<ExtensionEntry>> + Send + 'a {
async { None }
}
fn on_settle<'a>(
&'a self,
_ctx: &'a SettleContext<'a>,
) -> impl Future<Output = Option<ExtensionEntry>> + Send + 'a {
async { None }
}
}
pub trait DynExtension: Send + Sync {
fn id(&self) -> &'static str;
fn advertise(&self, ctx: &AdvertiseContext<'_>) -> Option<ExtensionEntry>;
fn on_verify<'a>(&'a self, ctx: &'a VerifyContext<'a>)
-> BoxFuture<'a, Option<ExtensionEntry>>;
fn on_settle<'a>(&'a self, ctx: &'a SettleContext<'a>)
-> BoxFuture<'a, Option<ExtensionEntry>>;
}
impl<T: Extension + ?Sized> DynExtension for T {
fn id(&self) -> &'static str {
<Self as Extension>::id(self)
}
fn advertise(&self, ctx: &AdvertiseContext<'_>) -> Option<ExtensionEntry> {
<Self as Extension>::advertise(self, ctx)
}
fn on_verify<'a>(
&'a self,
ctx: &'a VerifyContext<'a>,
) -> BoxFuture<'a, Option<ExtensionEntry>> {
Box::pin(<Self as Extension>::on_verify(self, ctx))
}
fn on_settle<'a>(
&'a self,
ctx: &'a SettleContext<'a>,
) -> BoxFuture<'a, Option<ExtensionEntry>> {
Box::pin(<Self as Extension>::on_settle(self, ctx))
}
}
#[derive(Clone, Default)]
pub struct ExtensionRegistry {
ordered: Vec<Arc<dyn DynExtension>>,
by_id: HashMap<CompactString, Arc<dyn DynExtension>>,
}
impl Debug for ExtensionRegistry {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
let ids: Vec<&str> = self.ordered.iter().map(|e| e.id()).collect();
f.debug_tuple("ExtensionRegistry").field(&ids).finish()
}
}
impl ExtensionRegistry {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register<E: Extension + 'static>(&mut self, extension: E) {
let id = extension.id();
let arc: Arc<dyn DynExtension> = Arc::new(extension);
self.ordered.retain(|existing| existing.id() != id);
self.ordered.push(Arc::clone(&arc));
let _ = self.by_id.insert(CompactString::from(id), arc);
}
#[must_use]
pub fn get(&self, id: &str) -> Option<&dyn DynExtension> {
self.by_id.get(id).map(AsRef::as_ref)
}
pub fn iter(&self) -> impl Iterator<Item = &dyn DynExtension> {
self.ordered.iter().map(AsRef::as_ref)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.ordered.is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.ordered.len()
}
#[must_use]
pub fn advertise(&self, ctx: &AdvertiseContext<'_>) -> Extensions {
let mut out = Extensions::new();
for ext in self.iter() {
if let Some(entry) = ext.advertise(ctx) {
out.insert(ext.id(), entry);
}
}
out
}
pub async fn collect_verify(&self, ctx: &VerifyContext<'_>) -> Extensions {
let mut out = Extensions::new();
for ext in self.iter() {
if let Some(entry) = ext.on_verify(ctx).await {
out.insert(ext.id(), entry);
}
}
out
}
pub async fn collect_settle(&self, ctx: &SettleContext<'_>) -> Extensions {
let mut out = Extensions::new();
for ext in self.iter() {
if let Some(entry) = ext.on_settle(ctx).await {
out.insert(ext.id(), entry);
}
}
out
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
struct StubExt(&'static str, serde_json::Value);
impl Extension for StubExt {
fn id(&self) -> &'static str {
self.0
}
fn advertise(&self, _: &AdvertiseContext<'_>) -> Option<ExtensionEntry> {
Some(ExtensionEntry::info(self.1.clone()))
}
}
#[test]
fn register_and_lookup() {
let mut registry = ExtensionRegistry::new();
registry.register(StubExt("bazaar", json!({"registered": true})));
registry.register(StubExt("other", json!({"x": 1})));
assert_eq!(registry.len(), 2);
assert!(registry.get("bazaar").is_some());
assert!(registry.get("missing").is_none());
}
#[test]
fn duplicate_registration_overwrites() {
let mut registry = ExtensionRegistry::new();
registry.register(StubExt("x", json!(1)));
registry.register(StubExt("x", json!(2)));
assert_eq!(registry.len(), 1);
let ctx = AdvertiseContext { requirement: None };
let ext = registry.advertise(&ctx);
let val = ext.get("x").unwrap().as_info().unwrap();
assert_eq!(val, &json!(2));
}
#[test]
fn advertise_emits_every_entry() {
let mut registry = ExtensionRegistry::new();
registry.register(StubExt("a", json!("A")));
registry.register(StubExt("b", json!("B")));
let ctx = AdvertiseContext { requirement: None };
let ext = registry.advertise(&ctx);
assert_eq!(ext.len(), 2);
}
}