use std::collections::{HashMap, HashSet};
use std::future::Future;
use std::pin::Pin;
use serde_json::Value;
use crate::dispatcher::{build_error_response, canonical_key, downcast_payload, RequestOrigin};
use crate::document::{ErrorResponse, TrustTask};
use crate::error::{ErrorPayload, RejectReason};
use crate::payload::Payload;
type HandlerFuture<R> = Pin<Box<dyn Future<Output = Result<R, RejectReason>> + Send>>;
type BoxedAsyncHandler<Ctx, R> =
Box<dyn Fn(TrustTask<Value>, Ctx) -> HandlerFuture<R> + Send + Sync>;
pub struct AsyncDispatcher<Ctx, R> {
handlers: HashMap<String, BoxedAsyncHandler<Ctx, R>>,
slugs: HashSet<String>,
}
impl<Ctx, R> Default for AsyncDispatcher<Ctx, R> {
fn default() -> Self {
Self::new()
}
}
impl<Ctx, R> AsyncDispatcher<Ctx, R> {
pub fn new() -> Self {
Self {
handlers: HashMap::new(),
slugs: HashSet::new(),
}
}
pub fn on_async<P, F, Fut>(mut self, handler: F) -> Self
where
P: Payload + 'static,
F: Fn(TrustTask<P>, Ctx) -> Fut + Send + Sync + 'static,
Fut: Future<Output = R> + Send + 'static,
R: Send + 'static,
{
let wrapped = move |doc: TrustTask<Value>, ctx: Ctx| -> HandlerFuture<R> {
let is_request = !doc.type_uri.is_response();
let typed = match downcast_payload::<P>(doc) {
Ok(typed) => typed,
Err(reason) => return Box::pin(std::future::ready(Err(reason))),
};
if is_request {
if let Err(reason) = typed.enforce_spec_policy() {
return Box::pin(std::future::ready(Err(reason)));
}
}
let fut = handler(typed, ctx);
Box::pin(async move { Ok(fut.await) })
};
self.slugs.insert(P::type_uri().slug().to_string());
self.handlers
.insert(canonical_key(&P::type_uri()), Box::new(wrapped));
self
}
pub async fn dispatch(&self, doc: TrustTask<Value>, ctx: Ctx) -> Result<R, RejectReason> {
let key = canonical_key(&doc.type_uri);
match self.handlers.get(&key) {
Some(handler) => handler(doc, ctx).await,
None if self.slugs.contains(doc.type_uri.slug()) => {
Err(RejectReason::UnsupportedVersion { type_uri: key })
}
None => Err(RejectReason::UnsupportedType { type_uri: key }),
}
}
#[allow(clippy::result_large_err)]
pub async fn dispatch_or_reject(
&self,
doc: TrustTask<Value>,
ctx: Ctx,
error_id: impl Into<String>,
) -> Result<R, ErrorResponse> {
let origin = RequestOrigin {
id: doc.id.clone(),
thread_id: doc.thread_id.clone(),
parent_thread_id: doc.parent_thread_id.clone(),
ceremony: doc.ceremony.clone(),
type_uri: doc.type_uri.to_string(),
issuer: doc.issuer.clone(),
recipient: doc.recipient.clone(),
};
match self.dispatch(doc, ctx).await {
Ok(value) => Ok(value),
Err(reason) => Err(build_error_response(
error_id.into(),
origin,
ErrorPayload::from(reason),
)),
}
}
pub fn registered_uris(&self) -> Vec<&str> {
let mut v: Vec<&str> = self.handlers.keys().map(String::as_str).collect();
v.sort_unstable();
v
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::specs::acl::grant::v0_1 as grant;
use crate::StandardCode;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
fn payload() -> grant::Payload {
grant::Payload {
entry: grant::AclEntry {
allowed_keys: None,
subject: "did:web:alice.example".into(),
role: "admin".into(),
scopes: vec![],
label: None,
created_at: None,
created_by: None,
updated_at: None,
updated_by: None,
expires_at: None,
approve: None,
step_up: None,
ext: None,
},
reason: None,
ext: None,
}
}
fn inbound(type_uri: &str, with_proof: bool) -> TrustTask<Value> {
let mut doc = TrustTask::for_payload("req-1", payload());
doc.issuer = Some("did:web:org.example".into());
doc.recipient = Some("did:web:maintainer.example".into());
doc.issued_at = Some(chrono::Utc::now());
if with_proof {
doc.proof = Some(crate::Proof {
proof_type: "DataIntegrityProof".into(),
cryptosuite: "eddsa-jcs-2022".into(),
verification_method: "did:web:org.example#key-1".into(),
created: chrono::Utc::now(),
proof_purpose: "assertionMethod".into(),
proof_value: "z000".into(),
extra: Default::default(),
});
}
let mut value = serde_json::to_value(&doc).unwrap();
value["type"] = Value::String(type_uri.to_string());
serde_json::from_value(value).unwrap()
}
const GRANT_V0_1: &str = "https://trusttasks.org/spec/acl/grant/0.1";
#[tokio::test]
async fn async_handler_awaits_and_returns_a_response() {
let calls = Arc::new(AtomicUsize::new(0));
let dispatcher =
AsyncDispatcher::<Arc<AtomicUsize>, TrustTask<grant::Response>>::new()
.on_async::<grant::Payload, _, _>(|req, ctx: Arc<AtomicUsize>| async move {
tokio::task::yield_now().await;
ctx.fetch_add(1, Ordering::SeqCst);
req.respond_with(
"resp-1",
grant::Response {
entry: req.payload.entry.clone(),
ext: None,
},
)
});
let response = dispatcher
.dispatch(inbound(GRANT_V0_1, true), Arc::clone(&calls))
.await
.expect("handler ran");
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"context reached the handler"
);
assert!(response.type_uri.is_response());
assert_eq!(response.thread_id.as_deref(), Some("req-1"));
assert_eq!(response.payload.entry.subject, "did:web:alice.example");
assert_eq!(
response.issuer.as_deref(),
Some("did:web:maintainer.example")
);
assert_eq!(response.recipient.as_deref(), Some("did:web:org.example"));
}
#[tokio::test]
async fn unknown_slug_is_unsupported_type() {
let dispatcher = AsyncDispatcher::<(), ()>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async {});
let doc = inbound("https://trusttasks.org/spec/never-heard-of-it/1.0", true);
let err = dispatcher.dispatch(doc, ()).await.unwrap_err();
assert!(
matches!(err, RejectReason::UnsupportedType { .. }),
"expected UnsupportedType, got {err:?}"
);
}
#[tokio::test]
async fn known_slug_at_unknown_version_is_unsupported_version() {
let dispatcher = AsyncDispatcher::<(), ()>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async {});
let doc = inbound("https://trusttasks.org/spec/acl/grant/9.9", true);
let err = dispatcher.dispatch(doc, ()).await.unwrap_err();
assert!(
matches!(err, RejectReason::UnsupportedVersion { .. }),
"expected UnsupportedVersion, got {err:?}"
);
}
#[tokio::test]
async fn dispatch_or_reject_carries_the_version_distinction_onto_the_wire() {
let dispatcher = AsyncDispatcher::<(), ()>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async {});
let err = dispatcher
.dispatch_or_reject(
inbound("https://trusttasks.org/spec/acl/grant/9.9", true),
(),
"err-1",
)
.await
.unwrap_err();
assert_eq!(err.payload.code, StandardCode::UnsupportedVersion.into());
assert_eq!(err.id, "err-1");
assert_eq!(err.thread_id.as_deref(), Some("req-1"));
assert_eq!(err.issuer.as_deref(), Some("did:web:maintainer.example"));
assert_eq!(err.recipient.as_deref(), Some("did:web:org.example"));
let err = dispatcher
.dispatch_or_reject(
inbound("https://trusttasks.org/spec/nope/1.0", true),
(),
"err-2",
)
.await
.unwrap_err();
assert_eq!(err.payload.code, StandardCode::UnsupportedType.into());
}
#[tokio::test]
async fn typed_spec_policy_runs_before_the_handler() {
let calls = Arc::new(AtomicUsize::new(0));
let dispatcher = AsyncDispatcher::<Arc<AtomicUsize>, ()>::new()
.on_async::<grant::Payload, _, _>(|_req, ctx: Arc<AtomicUsize>| async move {
ctx.fetch_add(1, Ordering::SeqCst);
});
let err = dispatcher
.dispatch(inbound(GRANT_V0_1, false), Arc::clone(&calls))
.await
.unwrap_err();
assert!(
matches!(err, RejectReason::ProofRequired),
"expected ProofRequired, got {err:?}"
);
assert_eq!(calls.load(Ordering::SeqCst), 0, "handler must not have run");
}
#[tokio::test]
async fn payload_that_does_not_match_is_malformed_request() {
let dispatcher = AsyncDispatcher::<(), ()>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async {});
let mut doc = inbound(GRANT_V0_1, true);
doc.payload = serde_json::json!({ "entry": "not an object" });
let err = dispatcher.dispatch(doc, ()).await.unwrap_err();
assert!(
matches!(err, RejectReason::MalformedRequest { .. }),
"expected MalformedRequest, got {err:?}"
);
}
#[tokio::test]
async fn request_fragment_and_bare_uri_route_together() {
let dispatcher = AsyncDispatcher::<(), &'static str>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async { "handled" });
assert_eq!(dispatcher.registered_uris(), vec![GRANT_V0_1]);
for uri in [
GRANT_V0_1,
"https://trusttasks.org/spec/acl/grant/0.1#request",
] {
assert_eq!(
dispatcher.dispatch(inbound(uri, true), ()).await.unwrap(),
"handled",
"{uri} should route to the registered handler"
);
}
}
#[tokio::test]
async fn dispatcher_is_shareable_across_tasks() {
let dispatcher = Arc::new(
AsyncDispatcher::<(), &'static str>::new()
.on_async::<grant::Payload, _, _>(|_req, _ctx| async { "handled" }),
);
let d = Arc::clone(&dispatcher);
let handle = tokio::spawn(async move { d.dispatch(inbound(GRANT_V0_1, true), ()).await });
assert_eq!(handle.await.unwrap().unwrap(), "handled");
}
}