use std::collections::{HashMap, VecDeque};
use tokio::sync::{Mutex, oneshot};
use tracing::{debug, warn};
struct PendingRequest {
nonce: Option<String>,
tx: oneshot::Sender<String>,
}
pub(crate) enum Routed {
Delivered,
Uncorrelated(String),
}
const PARKED_CAPACITY: usize = 256;
pub(crate) struct TspDemux {
pending: Mutex<HashMap<String, PendingRequest>>,
parked: Mutex<VecDeque<String>>,
}
impl TspDemux {
pub(crate) fn new() -> Self {
Self {
pending: Mutex::new(HashMap::new()),
parked: Mutex::new(VecDeque::new()),
}
}
pub(crate) async fn register(
&self,
request_id: String,
nonce: Option<String>,
) -> oneshot::Receiver<String> {
let (tx, rx) = oneshot::channel();
self.pending
.lock()
.await
.insert(request_id, PendingRequest { nonce, tx });
rx
}
pub(crate) async fn deregister(&self, request_id: &str) {
self.pending.lock().await.remove(request_id);
}
pub(crate) async fn route(&self, json: String) -> Routed {
let Ok(doc) = serde_json::from_str::<serde_json::Value>(&json) else {
return Routed::Uncorrelated(json);
};
let mut pending = self.pending.lock().await;
let hit = pending
.iter()
.find(|(id, p)| correlates(&doc, id, p.nonce.as_deref()))
.map(|(id, _)| id.clone());
match hit {
Some(id) => {
if let Some(p) = pending.remove(&id) {
let _ = p.tx.send(json);
}
Routed::Delivered
}
None => {
drop(pending);
debug!(
thread_id = doc
.get("threadId")
.and_then(|v| v.as_str())
.unwrap_or("<none>"),
"TSP frame matched no in-flight request (push, or a stale inbox entry)"
);
Routed::Uncorrelated(json)
}
}
}
pub(crate) async fn park(&self, doc: String) {
let mut parked = self.parked.lock().await;
if parked.len() >= PARKED_CAPACITY {
parked.pop_front();
warn!(
capacity = PARKED_CAPACITY,
"TSP push queue is full — discarding the oldest parked document. \
Nothing is draining inbound pushes on this session."
);
}
parked.push_back(doc);
}
pub(crate) async fn take_parked(&self) -> Option<String> {
self.parked.lock().await.pop_front()
}
pub(crate) async fn clear(&self) {
self.pending.lock().await.clear();
}
pub(crate) fn request_keys(document: &[u8]) -> Result<(String, Option<String>), String> {
let parsed: serde_json::Value = serde_json::from_slice(document)
.map_err(|e| format!("TSP request document is not JSON: {e}"))?;
let request_id = parsed
.get("id")
.and_then(|v| v.as_str())
.ok_or("TSP request document has no `id` to correlate its reply on")?
.to_string();
let nonce = parsed
.get("payload")
.and_then(|p| p.get("nonce"))
.and_then(|v| v.as_str())
.map(str::to_string);
Ok((request_id, nonce))
}
}
pub(crate) fn correlates(doc: &serde_json::Value, request_id: &str, nonce: Option<&str>) -> bool {
if doc.get("threadId").and_then(|v| v.as_str()) == Some(request_id) {
return true;
}
match nonce {
Some(n) => {
doc.get("payload")
.and_then(|p| p.get("nonce"))
.and_then(|v| v.as_str())
== Some(n)
}
None => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn response(thread_id: &str) -> String {
json!({
"id": "urn:uuid:reply",
"type": "https://trusttasks.org/spec/messaging/ping/0.1#response",
"threadId": thread_id,
"payload": { "ok": true },
})
.to_string()
}
#[test]
fn a_stale_frame_is_not_mistaken_for_our_reply() {
let stale = json!({
"id": "urn:uuid:reply-to-something-else",
"threadId": "urn:uuid:a-previous-run",
"payload": { "nonce": "some-other-nonce" },
});
assert!(!correlates(
&stale,
"urn:uuid:our-request",
Some("our-nonce")
));
}
#[test]
fn a_push_correlates_with_no_request() {
let push = json!({
"id": "urn:uuid:pushed",
"type": "https://trusttasks.org/spec/task-consent/request/0.1",
"payload": { "taskId": "t1" },
});
assert!(!correlates(
&push,
"urn:uuid:our-request",
Some("our-nonce")
));
}
#[test]
fn correlates_on_thread_id() {
let doc = json!({ "threadId": "urn:uuid:req-1" });
assert!(correlates(&doc, "urn:uuid:req-1", None));
assert!(!correlates(&doc, "urn:uuid:req-2", None));
}
#[test]
fn correlates_on_echoed_nonce_when_unthreaded() {
let doc = json!({ "id": "urn:uuid:fresh", "payload": { "nonce": "n-1" } });
assert!(correlates(&doc, "urn:uuid:req-1", Some("n-1")));
assert!(!correlates(&doc, "urn:uuid:req-1", Some("n-2")));
assert!(!correlates(&doc, "urn:uuid:req-1", None));
}
#[tokio::test]
async fn a_reply_reaches_its_own_waiter() {
let demux = TspDemux::new();
let mut rx = demux.register("urn:uuid:req-1".into(), None).await;
assert!(matches!(
demux.route(response("urn:uuid:req-1")).await,
Routed::Delivered
));
let got = rx.try_recv().expect("the waiter received its reply");
assert!(got.contains("urn:uuid:req-1"));
}
#[tokio::test]
async fn concurrent_waiters_do_not_steal_from_each_other() {
let demux = TspDemux::new();
let mut rx_a = demux.register("urn:uuid:a".into(), None).await;
let mut rx_b = demux.register("urn:uuid:b".into(), None).await;
demux.route(response("urn:uuid:b")).await;
assert!(
rx_a.try_recv().is_err(),
"A's waiter must not receive B's reply"
);
assert!(rx_b.try_recv().is_ok(), "B's waiter receives B's reply");
demux.route(response("urn:uuid:a")).await;
assert!(rx_a.try_recv().is_ok(), "A's waiter still receives its own");
}
#[tokio::test]
async fn an_uncorrelated_document_is_returned_and_can_be_parked() {
let demux = TspDemux::new();
let _rx = demux.register("urn:uuid:req-1".into(), None).await;
let push = json!({ "id": "urn:uuid:push", "type": "task-consent/request" }).to_string();
let Routed::Uncorrelated(doc) = demux.route(push.clone()).await else {
panic!("a push must not be delivered to an unrelated waiter");
};
assert_eq!(doc, push);
demux.park(doc).await;
assert_eq!(demux.take_parked().await.as_deref(), Some(push.as_str()));
assert!(demux.take_parked().await.is_none());
}
#[tokio::test]
async fn a_non_json_frame_is_uncorrelated_not_dropped() {
let demux = TspDemux::new();
assert!(matches!(
demux.route("not json at all".into()).await,
Routed::Uncorrelated(_)
));
}
#[tokio::test]
async fn clear_wakes_waiters_with_a_closed_channel() {
let demux = TspDemux::new();
let mut rx = demux.register("urn:uuid:req-1".into(), None).await;
demux.clear().await;
assert!(
matches!(rx.try_recv(), Err(oneshot::error::TryRecvError::Closed)),
"a cleared waiter must observe the channel closed, not stay pending"
);
}
#[tokio::test]
async fn parking_is_bounded_and_drops_the_oldest() {
let demux = TspDemux::new();
for i in 0..PARKED_CAPACITY + 1 {
demux.park(format!("doc-{i}")).await;
}
assert_eq!(
demux.take_parked().await.as_deref(),
Some("doc-1"),
"doc-0 must have been discarded to make room"
);
}
#[test]
fn request_keys_reads_id_and_nonce() {
let doc = json!({ "id": "urn:uuid:req-1", "payload": { "nonce": "n-1" } });
let (id, nonce) = TspDemux::request_keys(&serde_json::to_vec(&doc).unwrap()).unwrap();
assert_eq!(id, "urn:uuid:req-1");
assert_eq!(nonce.as_deref(), Some("n-1"));
}
#[test]
fn request_keys_refuses_a_document_with_no_id() {
let doc = json!({ "payload": { "nonce": "n-1" } });
assert!(TspDemux::request_keys(&serde_json::to_vec(&doc).unwrap()).is_err());
assert!(TspDemux::request_keys(b"not json").is_err());
}
}