use std::collections::HashMap;
use crate::{ApplicationFrame, Kind, SessionId, StreamKey};
pub(crate) struct RelayForward {
pub(crate) target: StreamKey,
pub(crate) frame: ApplicationFrame,
pub(crate) terminal: bool,
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::{
ApplicationError, ApplicationHead, BodyId, CorrelationId, ErrorCode, PROTOCOL_VERSION,
};
fn key(session: &str, corr: &str) -> StreamKey {
StreamKey {
session: SessionId::from(session),
corr: CorrelationId::from(corr),
}
}
fn frame(kind: Kind) -> ApplicationFrame {
ApplicationFrame {
head: ApplicationHead {
v: PROTOCOL_VERSION,
id: "frame".into(),
target: "target-node".into(),
subject: "opaque.subject".into(),
kind,
corr: Some("left-corr".into()),
seq: Some(7),
hops: Some(3),
path: vec!["owner".into()],
headers: serde_json::Map::from_iter([
("x-custom".into(), json!("value")),
("x-status".into(), json!(418)),
]),
error: (kind == Kind::Error).then(|| ApplicationError {
code: ErrorCode::Protocol,
message: "bad sibling".into(),
}),
},
body: (kind != Kind::Error && kind != Kind::Cancel).then(|| BodyId::from("body-1")),
}
}
fn paired() -> (RelayReducer, StreamKey, StreamKey) {
let left = key("client", "left-corr");
let right = key("peer", "right-corr");
let mut reducer = RelayReducer::default();
assert!(reducer.pair(left.clone(), right.clone()));
(reducer, left, right)
}
#[test]
fn unary_terminal_rewrites_and_retains_both_legs() {
let (mut reducer, left, right) = paired();
let original = frame(Kind::Response);
let forwarded = reducer.forward(&right, original.clone(), "relay").unwrap();
assert_eq!(forwarded.target, left);
assert_eq!(forwarded.frame.head.corr.as_deref(), Some("left-corr"));
assert_eq!(forwarded.frame.body, original.body);
assert_eq!(forwarded.frame.head.headers, original.head.headers);
assert!(forwarded.terminal);
assert_eq!(reducer.peer(&right), Some(&left));
assert_eq!(reducer.peer(&left), Some(&right));
}
#[test]
fn event_is_non_terminal_and_error_returns_with_annotated_path() {
let (mut reducer, left, right) = paired();
let original_event = frame(Kind::Event);
let event = reducer
.forward(&right, original_event.clone(), "relay")
.unwrap();
assert_eq!(event.target, left);
assert_eq!(event.frame, original_event);
assert!(!event.terminal);
assert_eq!(reducer.peer(&right), Some(&left));
let original_error = frame(Kind::Error);
let error = reducer
.forward(&right, original_error.clone(), "relay")
.unwrap();
assert_eq!(error.frame.head.id, original_error.head.id);
assert_eq!(error.frame.head.target, original_error.head.target);
assert_eq!(error.frame.head.subject, original_error.head.subject);
assert_eq!(error.frame.head.kind, original_error.head.kind);
assert_eq!(error.frame.head.seq, original_error.head.seq);
assert_eq!(error.frame.head.hops, original_error.head.hops);
assert_eq!(error.frame.head.headers, original_error.head.headers);
assert_eq!(error.frame.head.error, original_error.head.error);
assert_eq!(error.frame.head.path, ["owner", "relay"]);
assert!(error.terminal);
}
#[test]
fn response_preserves_opaque_payload_and_non_protocol_metadata() {
let (mut reducer, left, right) = paired();
let original = frame(Kind::Response);
let response = reducer.forward(&right, original.clone(), "relay").unwrap();
assert_eq!(response.target, left);
assert_eq!(response.frame, original);
}
#[test]
fn cancel_in_either_direction_rewrites_and_retires() {
for from_left in [true, false] {
let (mut reducer, left, right) = paired();
let source = if from_left { &left } else { &right };
let target = if from_left { &right } else { &left };
let cancel = reducer
.forward(source, frame(Kind::Cancel), "relay")
.unwrap();
assert_eq!(&cancel.target, target);
assert_eq!(
cancel.frame.head.corr.as_deref(),
Some(target.corr.as_str())
);
assert!(cancel.terminal);
assert_eq!(reducer.peer(source), Some(target));
}
}
#[test]
fn closing_either_session_removes_pair_and_identifies_peer_cleanup() {
let (mut reducer, left, right) = paired();
assert_eq!(
reducer.drain_session(&left.session),
vec![(right.clone(), false)]
);
assert_eq!(reducer.peer(&right), None);
let (mut reducer, left, right) = paired();
assert_eq!(
reducer.drain_session(&right.session),
vec![(left.clone(), true)]
);
assert_eq!(reducer.peer(&left), None);
}
}
struct RelayLeg {
peer: StreamKey,
annotate_error: bool,
}
#[derive(Default)]
pub(crate) struct RelayReducer {
pairs: HashMap<StreamKey, RelayLeg>,
}
impl RelayReducer {
pub(crate) fn pair(&mut self, left: StreamKey, right: StreamKey) -> bool {
if self.pairs.contains_key(&left) || self.pairs.contains_key(&right) {
return false;
}
self.pairs.insert(
left.clone(),
RelayLeg {
peer: right.clone(),
annotate_error: false,
},
);
self.pairs.insert(
right,
RelayLeg {
peer: left,
annotate_error: true,
},
);
true
}
pub(crate) fn peer(&self, stream: &StreamKey) -> Option<&StreamKey> {
self.pairs.get(stream).map(|leg| &leg.peer)
}
pub(crate) fn forward(
&mut self,
stream: &StreamKey,
mut frame: ApplicationFrame,
node: &str,
) -> Option<RelayForward> {
let leg = self.pairs.get(stream)?;
let target = leg.peer.clone();
let annotate_error = leg.annotate_error;
let terminal = matches!(frame.head.kind, Kind::Response | Kind::Error | Kind::Cancel);
frame.head.corr = Some(target.corr.as_str().to_string());
if frame.head.kind == Kind::Error && annotate_error {
frame.head.path.push(node.to_string());
}
Some(RelayForward {
target,
frame,
terminal,
})
}
pub(crate) fn drain_session(&mut self, session: &SessionId) -> Vec<(StreamKey, bool)> {
let streams: Vec<_> = self
.pairs
.keys()
.filter(|stream| &stream.session == session)
.cloned()
.collect();
streams
.into_iter()
.filter_map(|stream| {
let leg = self.pairs.get(&stream)?;
let peer = leg.peer.clone();
let peer_is_source = leg.annotate_error;
self.remove(&stream);
Some((peer, peer_is_source))
})
.collect()
}
pub(crate) fn remove_pair(&mut self, stream: &StreamKey) -> Option<(StreamKey, bool)> {
let leg = self.pairs.remove(stream)?;
self.pairs.remove(&leg.peer);
Some((leg.peer, leg.annotate_error))
}
fn remove(&mut self, stream: &StreamKey) {
if let Some(leg) = self.pairs.remove(stream) {
self.pairs.remove(&leg.peer);
}
}
}