unb-core 2.0.3

Core unb protocol types: envelope, session, routing, taxonomy
Documentation
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);
        }
    }
}