use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use async_trait::async_trait;
use serde_json::Value;
use tokio::sync::oneshot;
use trql_client::{TransportKind, TrqlError};
use verify_trust::RegistryChannel;
use crate::transport::{InboundDoc, Via, VtcLink};
const REPLY_TIMEOUT: Duration = Duration::from_secs(30);
type Waiter = (Via, oneshot::Sender<Value>);
#[derive(Debug)]
pub struct RegistryReplies {
registry_did: String,
pending: Mutex<HashMap<String, Waiter>>,
}
impl RegistryReplies {
pub fn new(registry_did: impl Into<String>) -> Self {
RegistryReplies {
registry_did: registry_did.into(),
pending: Mutex::new(HashMap::new()),
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<String, Waiter>> {
self.pending.lock().unwrap_or_else(|p| p.into_inner())
}
fn register(&self, id: &str, via: Via) -> (Pending<'_>, oneshot::Receiver<Value>) {
let (tx, rx) = oneshot::channel();
self.lock().insert(id.to_string(), (via, tx));
(
Pending {
replies: self,
id: id.to_string(),
},
rx,
)
}
#[cfg(test)]
fn in_flight(&self) -> usize {
self.lock().len()
}
pub fn route(&self, inbound: &InboundDoc) -> bool {
let Some(sender) = inbound.authenticated_sender.as_deref() else {
return false;
};
let sender = sender.split_once('#').map_or(sender, |(did, _)| did);
if sender != self.registry_did {
return false;
}
let Some(thread) = inbound.doc.get("threadId").and_then(Value::as_str) else {
return false;
};
let waiter = {
let mut pending = self.lock();
if pending
.get(thread)
.is_none_or(|(via, _)| *via != inbound.via)
{
return false;
}
match pending.remove(thread) {
Some((_, waiter)) => waiter,
None => return false,
}
};
let _ = waiter.send(inbound.doc.clone());
true
}
}
struct Pending<'a> {
replies: &'a RegistryReplies,
id: String,
}
impl Drop for Pending<'_> {
fn drop(&mut self) {
self.replies.lock().remove(&self.id);
}
}
pub struct BridgeRegistryChannel {
link: Arc<dyn VtcLink>,
did: String,
replies: Arc<RegistryReplies>,
via: Via,
timeout: Duration,
}
impl BridgeRegistryChannel {
pub fn new(
link: Arc<dyn VtcLink>,
did: impl Into<String>,
replies: Arc<RegistryReplies>,
via: Via,
) -> Self {
BridgeRegistryChannel {
link,
did: did.into(),
replies,
via,
timeout: REPLY_TIMEOUT,
}
}
#[doc(hidden)]
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
}
#[async_trait]
impl RegistryChannel for BridgeRegistryChannel {
fn kind(&self) -> TransportKind {
match self.via {
Via::Tsp => TransportKind::Tsp,
Via::Didcomm => TransportKind::Didcomm,
}
}
fn sender_did(&self) -> &str {
&self.did
}
async fn exchange(&self, recipient: &str, request: Value) -> Result<Value, TrqlError> {
let kind = self.kind();
let transport = |detail: String| TrqlError::Transport { kind, detail };
if recipient != self.replies.registry_did {
return Err(TrqlError::Config(format!(
"registry channel is for {}, not {recipient}",
self.replies.registry_did
)));
}
let id = request
.get("id")
.and_then(Value::as_str)
.ok_or_else(|| TrqlError::Contract("request document has no id".to_string()))?
.to_string();
let (_pending, reply) = self.replies.register(&id, self.via);
if let Err(e) = self.link.send_via(recipient, &request, self.via).await {
return Err(transport(format!("sending to the registry: {e:#}")));
}
match tokio::time::timeout(self.timeout, reply).await {
Ok(Ok(doc)) => Ok(doc),
Ok(Err(_)) => Err(transport("the reply channel closed".to_string())),
Err(_) => Err(TrqlError::Timeout {
kind,
waited_secs: self.timeout.as_secs(),
}),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transport::memory::ChannelLink;
const REGISTRY: &str = "did:webvh:QmReg:registry.example";
fn inbound(sender: Option<&str>, doc: Value) -> InboundDoc {
InboundDoc {
doc,
authenticated_sender: sender.map(str::to_string),
via: Via::Didcomm,
}
}
#[tokio::test]
async fn a_query_goes_out_as_the_bridge_and_its_proven_answer_comes_back() {
let (link, mut sent) = ChannelLink::new();
let replies = Arc::new(RegistryReplies::new(REGISTRY));
let channel = BridgeRegistryChannel::new(
Arc::new(link),
"did:key:z6MkBridge",
replies.clone(),
Via::Didcomm,
);
let answer = tokio::spawn(async move {
channel
.exchange(
REGISTRY,
serde_json::json!({ "id": "urn:uuid:q", "issuer": "did:key:z6MkBridge" }),
)
.await
});
let (to, doc) = sent.recv().await.unwrap();
assert_eq!(to, REGISTRY);
assert_eq!(doc["issuer"], "did:key:z6MkBridge");
let forged =
serde_json::json!({ "threadId": "urn:uuid:q", "payload": { "authorized": true } });
assert!(!replies.route(&inbound(Some("did:key:z6MkMallory"), forged.clone())));
assert!(!replies.route(&inbound(None, forged)));
let real =
serde_json::json!({ "threadId": "urn:uuid:q", "payload": { "authorized": false } });
assert!(replies.route(&inbound(Some(&format!("{REGISTRY}#key-2")), real)));
let got = answer.await.unwrap().unwrap();
assert_eq!(got["payload"]["authorized"], false);
}
#[tokio::test]
async fn a_tsp_query_goes_out_over_tsp_and_only_a_tsp_answer_is_taken() {
let (link, mut sent) = ChannelLink::new();
let link = Arc::new(link);
let replies = Arc::new(RegistryReplies::new(REGISTRY));
let channel = BridgeRegistryChannel::new(
link.clone(),
"did:key:z6MkBridge",
replies.clone(),
Via::Tsp,
);
assert_eq!(channel.kind(), TransportKind::Tsp);
let answer = tokio::spawn(async move {
channel
.exchange(REGISTRY, serde_json::json!({ "id": "urn:uuid:q" }))
.await
});
sent.recv().await.unwrap();
assert_eq!(
link.pinned(),
[Via::Tsp],
"pinned to TSP, not the link's choice"
);
let doc =
serde_json::json!({ "threadId": "urn:uuid:q", "payload": { "authorized": true } });
assert!(!replies.route(&inbound(Some(REGISTRY), doc.clone())));
let mallory = InboundDoc {
via: Via::Tsp,
..inbound(Some("did:key:z6MkMallory"), doc.clone())
};
assert!(!replies.route(&mallory));
let real = InboundDoc {
via: Via::Tsp,
..inbound(Some(REGISTRY), doc)
};
assert!(replies.route(&real));
assert_eq!(
answer.await.unwrap().unwrap()["payload"]["authorized"],
true
);
}
#[tokio::test]
async fn a_registry_that_never_answers_times_out_and_later_mail_is_not_taken() {
let (link, _sent) = ChannelLink::new();
let replies = Arc::new(RegistryReplies::new(REGISTRY));
let channel = BridgeRegistryChannel::new(
Arc::new(link),
"did:key:z6MkBridge",
replies.clone(),
Via::Didcomm,
)
.with_timeout(Duration::from_millis(20));
let e = channel
.exchange(REGISTRY, serde_json::json!({ "id": "urn:uuid:q" }))
.await
.unwrap_err();
assert!(matches!(e, TrqlError::Timeout { .. }), "{e}");
let late = serde_json::json!({ "threadId": "urn:uuid:q" });
assert!(!replies.route(&inbound(Some(REGISTRY), late)));
}
#[tokio::test]
async fn an_exchange_dropped_mid_wait_leaves_no_query_in_flight() {
let (link, mut sent) = ChannelLink::new();
let replies = Arc::new(RegistryReplies::new(REGISTRY));
let channel = BridgeRegistryChannel::new(
Arc::new(link),
"did:key:z6MkBridge",
replies.clone(),
Via::Didcomm,
);
let waiting = tokio::spawn(async move {
channel
.exchange(REGISTRY, serde_json::json!({ "id": "urn:uuid:q" }))
.await
});
sent.recv().await.unwrap();
assert_eq!(replies.in_flight(), 1);
waiting.abort();
assert!(waiting.await.unwrap_err().is_cancelled());
assert_eq!(replies.in_flight(), 0, "an aborted query must unregister");
let late = serde_json::json!({ "threadId": "urn:uuid:q" });
assert!(!replies.route(&inbound(Some(REGISTRY), late)));
}
#[tokio::test]
async fn a_link_that_is_down_fails_the_query_rather_than_waiting() {
let link = crate::transport::SupervisedLink::new(); let replies = Arc::new(RegistryReplies::new(REGISTRY));
let channel =
BridgeRegistryChannel::new(Arc::new(link), "did:key:z6MkBridge", replies, Via::Didcomm);
let e = channel
.exchange(REGISTRY, serde_json::json!({ "id": "urn:uuid:q" }))
.await
.unwrap_err();
assert!(matches!(e, TrqlError::Transport { .. }), "{e}");
}
}