#![cfg(feature = "didcomm")]
use std::sync::Arc;
use std::time::Duration;
use affinidi_messaging_core::MessageTransport;
use affinidi_tdk::messaging::DidCommTransport;
use affinidi_tdk::messaging::profiles::ATMProfile;
use async_trait::async_trait;
use serde_json::json;
use crate::didcomm_bridge::DIDCommBridge;
use crate::messaging::handshake::{
HandshakeStage, ListenerProver, ProverFailure, ResolvedMediator,
};
const TRUST_PING_TYPE: &str = "https://didcomm.org/trust-ping/2.0/ping";
const TRUST_PING_RESPONSE_TYPE: &str = "https://didcomm.org/trust-ping/2.0/ping-response";
const PROBLEM_REPORT_TYPE: &str = "https://didcomm.org/report-problem/2.0/problem-report";
pub struct DIDCommServiceProver {
bridge: Arc<DIDCommBridge>,
#[allow(dead_code)]
vta_did: String,
}
impl DIDCommServiceProver {
pub fn new(bridge: Arc<DIDCommBridge>, vta_did: impl Into<String>) -> Self {
Self {
bridge,
vta_did: vta_did.into(),
}
}
}
#[async_trait]
impl ListenerProver for DIDCommServiceProver {
async fn prove(
&self,
resolved: &ResolvedMediator,
vta_did: &str,
timeout: Duration,
) -> Result<(), ProverFailure> {
let service = self
.bridge
.messaging_handle()
.ok_or_else(|| ProverFailure {
stage: HandshakeStage::Connect,
cause: "delivery-layer messaging service is not running".to_string(),
})?;
let atm = self.bridge.atm().ok_or_else(|| ProverFailure {
stage: HandshakeStage::Connect,
cause: "ATM unavailable (messaging not started)".to_string(),
})?;
let candidate_id = resolved.mediator_did.clone();
let profile = match ATMProfile::new(
&atm,
Some(candidate_id.clone()),
vta_did.to_string(),
Some(candidate_id.clone()),
)
.await
{
Ok(p) => p,
Err(e) => {
return Err(ProverFailure {
stage: HandshakeStage::Connect,
cause: format!("create candidate profile: {e}"),
});
}
};
let profile = match atm.profile_add(&profile, false).await {
Ok(p) => p,
Err(e) => {
return Err(ProverFailure {
stage: HandshakeStage::Connect,
cause: format!("register candidate profile: {e}"),
});
}
};
let outcome = self
.connect_and_ping(&atm, &service, &profile, &candidate_id, vta_did, timeout)
.await;
if let Err(failure) = outcome {
service.remove_transport(&candidate_id);
drop_candidate_socket(&atm, &candidate_id).await;
return Err(failure);
}
Ok(())
}
}
impl DIDCommServiceProver {
async fn connect_and_ping(
&self,
atm: &affinidi_tdk::messaging::ATM,
service: &Arc<affinidi_messaging_delivery::MessagingService>,
profile: &Arc<ATMProfile>,
candidate_id: &str,
vta_did: &str,
timeout: Duration,
) -> Result<(), ProverFailure> {
let connect_fail = |cause: String| ProverFailure {
stage: HandshakeStage::Connect,
cause,
};
match tokio::time::timeout(timeout, atm.profile_enable_websocket(profile)).await {
Ok(Ok(())) => {}
Ok(Err(e)) => return Err(connect_fail(format!("enable candidate websocket: {e}"))),
Err(_) => {
return Err(connect_fail(
"timeout enabling candidate mediator websocket".to_string(),
));
}
}
let transport: Arc<dyn MessageTransport> =
match DidCommTransport::new(atm.clone(), profile.clone()).await {
Ok(t) => Arc::new(t),
Err(e) => {
return Err(connect_fail(format!(
"bind candidate DidComm transport: {e}"
)));
}
};
service.add_transport(candidate_id.to_string(), transport);
self.bridge
.send_and_wait_via(
candidate_id,
vta_did, TRUST_PING_TYPE,
json!({ "response_requested": true }),
TRUST_PING_RESPONSE_TYPE,
PROBLEM_REPORT_TYPE,
timeout.as_secs(),
)
.await
.map_err(|e| ProverFailure {
stage: HandshakeStage::TrustPing,
cause: format!("trust-ping round-trip failed: {e}"),
})?;
Ok(())
}
}
async fn drop_candidate_socket(atm: &affinidi_tdk::messaging::ATM, candidate_id: &str) {
if let Err(e) = atm.profile_remove(candidate_id).await {
tracing::warn!(
candidate = %candidate_id,
error = %e,
"could not stop the rejected candidate's websocket"
);
}
}
pub async fn try_build_from_parts(
bridge: &Arc<DIDCommBridge>,
vta_did: &str,
_secrets_resolver: &Arc<affinidi_tdk::secrets_resolver::ThreadedSecretsResolver>,
_signing_vm_id: &str,
_ka_vm_id: &str,
) -> Option<DIDCommServiceProver> {
bridge.messaging_handle()?;
Some(DIDCommServiceProver::new(Arc::clone(bridge), vta_did))
}
#[cfg(test)]
mod tests {
use crate::messaging::handshake::HandshakeStage;
#[test]
fn handshake_stages_used_by_prover() {
let stages = [HandshakeStage::Connect, HandshakeStage::TrustPing];
assert_eq!(stages.len(), 2);
}
}