use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use vta_sdk::error::VtaError;
pub const REPLY_TIMEOUT_BREAKER: u32 = 2;
const TSP_REPLY_TIMEOUT_PREFIX: &str = "timed out waiting for the TSP reply";
const DIDCOMM_REPLY_TIMEOUT: &str = "timeout waiting for DIDComm response";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReplyOutcome {
Replied,
TimedOut,
Unknown,
}
#[must_use]
pub fn is_reply_timeout(e: &VtaError) -> bool {
match e {
VtaError::TspTransport(msg) => msg.starts_with(TSP_REPLY_TIMEOUT_PREFIX),
VtaError::DidcommTransport(msg) => msg.contains(DIDCOMM_REPLY_TIMEOUT),
_ => false,
}
}
#[must_use]
pub fn classify<T>(outcome: &Result<T, VtaError>) -> ReplyOutcome {
match outcome {
Ok(_) => ReplyOutcome::Replied,
Err(e) if is_reply_timeout(e) => ReplyOutcome::TimedOut,
Err(
VtaError::Auth(_)
| VtaError::NotFound(_)
| VtaError::Validation(_)
| VtaError::Forbidden(_)
| VtaError::Conflict(_)
| VtaError::Gone(_)
| VtaError::Server { .. }
| VtaError::DidcommRemote { .. }
| VtaError::ConsentRequired { .. }
| VtaError::LastServiceRefused
| VtaError::ServiceNotPresent
| VtaError::ServiceAlreadyEnabled
| VtaError::DrainTtlOutOfBounds { .. }
| VtaError::NoPriorMutation
| VtaError::NoMatchingProtocol { .. }
| VtaError::UnsupportedTaskType { .. }
| VtaError::Unavailable { .. }
| VtaError::RateLimited { .. },
) => ReplyOutcome::Replied,
Err(_) => ReplyOutcome::Unknown,
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct ReceiveHealth {
pub consecutive_reply_timeouts: u32,
pub since_last_reply: Option<Duration>,
pub since_last_timeout: Option<Duration>,
}
impl ReceiveHealth {
#[must_use]
pub fn replies_not_arriving(&self) -> bool {
self.consecutive_reply_timeouts >= REPLY_TIMEOUT_BREAKER
}
}
#[derive(Debug, Default)]
struct Leg {
consecutive_timeouts: u32,
last_reply: Option<Instant>,
last_timeout: Option<Instant>,
}
#[derive(Clone, Debug, Default)]
pub struct ReceiveLegTracker {
leg: Arc<Mutex<Leg>>,
}
impl ReceiveLegTracker {
#[must_use]
pub fn new() -> Self {
Self::default()
}
fn with<R>(&self, f: impl FnOnce(&mut Leg) -> R) -> R {
let mut leg = self
.leg
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
f(&mut leg)
}
pub fn observe<T>(&self, outcome: Result<T, VtaError>) -> Result<T, VtaError> {
self.record(classify(&outcome), Instant::now());
outcome
}
pub fn record(&self, outcome: ReplyOutcome, now: Instant) {
self.with(|leg| match outcome {
ReplyOutcome::Replied => {
leg.consecutive_timeouts = 0;
leg.last_reply = Some(now);
}
ReplyOutcome::TimedOut => {
leg.consecutive_timeouts = leg.consecutive_timeouts.saturating_add(1);
leg.last_timeout = Some(now);
}
ReplyOutcome::Unknown => {}
});
}
pub fn reset(&self) {
self.with(|leg| *leg = Leg::default());
}
#[must_use]
pub fn health_at(&self, now: Instant) -> ReceiveHealth {
self.with(|leg| ReceiveHealth {
consecutive_reply_timeouts: leg.consecutive_timeouts,
since_last_reply: leg.last_reply.map(|t| now.saturating_duration_since(t)),
since_last_timeout: leg.last_timeout.map(|t| now.saturating_duration_since(t)),
})
}
#[must_use]
pub fn health(&self) -> ReceiveHealth {
self.health_at(Instant::now())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tsp_timeout() -> VtaError {
VtaError::TspTransport(format!(
"{TSP_REPLY_TIMEOUT_PREFIX} to request 'urn:uuid:1'"
))
}
fn didcomm_timeout() -> VtaError {
VtaError::DidcommTransport(DIDCOMM_REPLY_TIMEOUT.into())
}
#[test]
fn reply_timeouts_are_recognised_on_both_transports() {
assert!(is_reply_timeout(&tsp_timeout()));
assert!(is_reply_timeout(&didcomm_timeout()));
assert_eq!(classify::<()>(&Err(tsp_timeout())), ReplyOutcome::TimedOut);
assert_eq!(
classify::<()>(&Err(didcomm_timeout())),
ReplyOutcome::TimedOut
);
}
#[test]
fn a_failure_to_send_says_nothing_about_replies() {
for e in [
VtaError::TspTransport("failed to seal frame".into()),
VtaError::DidcommTransport("message pickup error: socket closed".into()),
VtaError::Protocol("bad shape".into()),
VtaError::Other("x".into()),
] {
assert!(!is_reply_timeout(&e));
assert_eq!(classify::<()>(&Err(e)), ReplyOutcome::Unknown);
}
}
#[test]
fn a_refusal_from_the_vta_is_a_reply() {
for e in [
VtaError::NotFound("x".into()),
VtaError::Auth("expired".into()),
VtaError::Server {
status: 500,
body: String::new(),
},
] {
assert_eq!(classify::<()>(&Err(e)), ReplyOutcome::Replied);
}
assert_eq!(classify::<u8>(&Ok(1)), ReplyOutcome::Replied);
}
#[test]
fn two_timeouts_in_a_row_mean_replies_are_not_arriving() {
let t = ReceiveLegTracker::new();
assert_eq!(t.health(), ReceiveHealth::default());
let _ = t.observe::<()>(Err(tsp_timeout()));
assert!(
!t.health().replies_not_arriving(),
"one is below the breaker"
);
let _ = t.observe::<()>(Err(VtaError::TspTransport("failed to seal frame".into())));
let _ = t.observe::<()>(Err(didcomm_timeout()));
let h = t.health();
assert_eq!(h.consecutive_reply_timeouts, 2);
assert!(h.replies_not_arriving());
assert!(h.since_last_timeout.is_some());
}
#[test]
fn any_reply_resets_the_count() {
let t = ReceiveLegTracker::new();
let _ = t.observe::<()>(Err(tsp_timeout()));
let _ = t.observe::<()>(Err(tsp_timeout()));
let _ = t.observe::<()>(Err(VtaError::NotFound("no such device".into())));
let h = t.health();
assert_eq!(h.consecutive_reply_timeouts, 0);
assert!(h.since_last_reply.is_some());
}
#[test]
fn clones_share_one_leg_and_reset_starts_afresh() {
let t = ReceiveLegTracker::new();
let other = t.clone();
let _ = other.observe::<()>(Err(tsp_timeout()));
let _ = t.observe::<()>(Err(tsp_timeout()));
assert!(t.health().replies_not_arriving());
t.reset();
assert_eq!(other.health(), ReceiveHealth::default());
}
#[test]
fn observe_hands_the_outcome_back_unchanged() {
let t = ReceiveLegTracker::new();
assert_eq!(t.observe::<u8>(Ok(7)).ok(), Some(7));
assert!(matches!(
t.observe::<u8>(Err(VtaError::NotFound("x".into()))),
Err(VtaError::NotFound(_))
));
}
}