use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use serde_json::Value;
use tokio::sync::oneshot;
use trust_tasks_rs::TrustTask;
#[must_use]
pub fn reply_thread_of(request: &Value) -> Option<&str> {
request
.get("threadId")
.and_then(Value::as_str)
.or_else(|| request.get("id").and_then(Value::as_str))
}
#[derive(Clone, Default)]
pub struct PendingReplies {
inner: Arc<Mutex<HashMap<String, oneshot::Sender<TrustTask<Value>>>>>,
}
impl PendingReplies {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn register(&self, thread: &str) -> oneshot::Receiver<TrustTask<Value>> {
let (tx, rx) = oneshot::channel();
self.lock().insert(thread.to_string(), tx);
rx
}
pub fn abandon(&self, thread: &str) {
self.lock().remove(thread);
}
pub fn complete(&self, document: &TrustTask<Value>) -> bool {
let Some(thread_id) = document.thread_id.as_deref() else {
return false;
};
let Some(waiter) = self.lock().remove(thread_id) else {
return false;
};
let _ = waiter.send(document.clone());
true
}
#[must_use]
pub fn outstanding(&self) -> usize {
self.lock().len()
}
fn lock(
&self,
) -> std::sync::MutexGuard<'_, HashMap<String, oneshot::Sender<TrustTask<Value>>>> {
self.inner.lock().unwrap_or_else(|e| e.into_inner())
}
}
#[cfg(test)]
mod tests {
use super::*;
use trust_tasks_rs::TypeUri;
fn request(id: &str, thread: Option<&str>) -> TrustTask<Value> {
let type_uri: TypeUri = "https://trusttasks.org/spec/auth/revoke-session/0.1"
.parse()
.expect("a well-formed Type URI");
let mut doc = TrustTask::new(id, type_uri, serde_json::json!({}));
doc.thread_id = thread.map(str::to_string);
doc
}
#[test]
fn the_key_matches_the_reply_the_framework_builds() {
let fresh = request("urn:uuid:req-1", None);
let answer = fresh.respond_with("urn:uuid:res-1", serde_json::json!({}));
assert_eq!(
reply_thread_of(&serde_json::to_value(&fresh).unwrap()),
answer.thread_id.as_deref(),
"a request with no threadId is answered in a thread named by its id"
);
let threaded = request("urn:uuid:req-2", Some("urn:uuid:thread-a"));
let answer = threaded.respond_with("urn:uuid:res-2", serde_json::json!({}));
assert_eq!(answer.thread_id.as_deref(), Some("urn:uuid:thread-a"));
assert_eq!(
reply_thread_of(&serde_json::to_value(&threaded).unwrap()),
answer.thread_id.as_deref(),
"a request already in a thread is answered in that thread"
);
}
#[test]
fn the_key_matches_a_rejection_too() {
let threaded = request("urn:uuid:req-3", Some("urn:uuid:thread-b"));
let reject = threaded.reject_with(
"urn:uuid:err-1",
trust_tasks_rs::RejectReason::MalformedRequest {
reason: "nope".into(),
},
);
assert_eq!(
reply_thread_of(&serde_json::to_value(&threaded).unwrap()),
reject.thread_id.as_deref(),
"a rejection threads the same way a success does"
);
}
#[tokio::test]
async fn a_reply_reaches_the_waiter_and_is_not_dispatched() {
let replies = PendingReplies::new();
let waiting = replies.register("urn:uuid:thread-c");
assert_eq!(replies.outstanding(), 1);
let reply = request("urn:uuid:res-4", Some("urn:uuid:thread-c"));
assert!(
replies.complete(&reply),
"`true` is what tells the spine not to dispatch this as a request"
);
let received = waiting.await.expect("the waiter is woken");
assert_eq!(received.id, "urn:uuid:res-4");
assert_eq!(
replies.outstanding(),
0,
"a delivered waiter is removed, so a duplicate cannot be delivered twice"
);
}
#[test]
fn a_document_nobody_is_waiting_for_falls_through() {
let replies = PendingReplies::new();
let _waiting = replies.register("urn:uuid:thread-d");
let other = request("urn:uuid:req-5", Some("urn:uuid:thread-elsewhere"));
assert!(!replies.complete(&other));
let opening = request("urn:uuid:thread-d", None);
assert!(
!replies.complete(&opening),
"a request whose id collides with an outstanding thread is still a request"
);
assert_eq!(replies.outstanding(), 1, "neither took the waiter");
}
#[test]
fn an_abandoned_waiter_lets_a_late_reply_fall_through() {
let replies = PendingReplies::new();
let _waiting = replies.register("urn:uuid:thread-e");
replies.abandon("urn:uuid:thread-e");
assert_eq!(replies.outstanding(), 0);
let late = request("urn:uuid:res-6", Some("urn:uuid:thread-e"));
assert!(
!replies.complete(&late),
"after a timeout the entry is gone, so a late answer is not claimed"
);
}
#[test]
fn a_reply_whose_waiter_gave_up_is_still_claimed() {
let replies = PendingReplies::new();
drop(replies.register("urn:uuid:thread-f"));
let reply = request("urn:uuid:res-7", Some("urn:uuid:thread-f"));
assert!(replies.complete(&reply));
assert_eq!(replies.outstanding(), 0);
}
}