use std::fmt;
use crate::header::{GuaranteeSet, Hello};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Agreed {
pub version: u64,
pub send_max_header_bytes: u64,
pub guarantees: GuaranteeSet,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum NegotiateError {
NoCommonVersion,
UnsupportedRequiredCapability(u64),
IncomparableGuarantee {
dimension: &'static str,
},
GuaranteeNotOffered {
peers_requirement: bool,
},
}
impl fmt::Display for NegotiateError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
NegotiateError::NoCommonVersion => {
f.write_str("no wire protocol version is supported by both peers")
}
NegotiateError::UnsupportedRequiredCapability(c) => {
write!(f, "peer requires unsupported capability {c}")
}
NegotiateError::IncomparableGuarantee { dimension } => {
write!(
f,
"the two peers declare different {dimension}, which is not an ordered dimension"
)
}
NegotiateError::GuaranteeNotOffered { peers_requirement } => {
let who = if *peers_requirement { "peer" } else { "local" };
write!(
f,
"the effective guarantee set does not reach the {who} requirement"
)
}
}
}
}
impl std::error::Error for NegotiateError {}
impl From<NegotiateError> for weida_core::Error {
fn from(e: NegotiateError) -> weida_core::Error {
weida_core::Error::Negotiation(e.to_string())
}
}
pub fn negotiate(ours: &Hello, theirs: &Hello) -> Result<Agreed, NegotiateError> {
let version = theirs
.versions
.iter()
.copied()
.filter(|v| ours.versions.contains(v))
.max()
.ok_or(NegotiateError::NoCommonVersion)?;
if let Some(missing) = theirs
.required_capabilities
.iter()
.copied()
.find(|c| !ours.capabilities.contains(c))
{
return Err(NegotiateError::UnsupportedRequiredCapability(missing));
}
let guarantees = ours
.offered()
.intersect(&theirs.offered())
.map_err(|dimension| NegotiateError::IncomparableGuarantee { dimension })?;
for (required, peers_requirement) in [(theirs.required(), true), (ours.required(), false)] {
if !guarantees.reaches(&required) {
return Err(NegotiateError::GuaranteeNotOffered { peers_requirement });
}
}
Ok(Agreed {
version,
send_max_header_bytes: theirs.max_header_bytes,
guarantees,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::header::{
Acknowledgement, Backpressure, Deduplication, Durability, OrderingMode, ProducerNaming,
};
fn hello(versions: &[u64], caps: &[u64], required: &[u64], max_header: u64) -> Hello {
Hello {
versions: versions.to_vec(),
max_header_bytes: max_header,
max_transfers: 1024,
capabilities: caps.to_vec(),
required_capabilities: required.to_vec(),
guarantees_offered: None,
guarantees_required: None,
}
}
fn declaring(offered: GuaranteeSet, required: GuaranteeSet) -> Hello {
Hello {
guarantees_offered: Some(offered),
guarantees_required: Some(required),
..Hello::v0(16384, 1024)
}
}
#[test]
fn v0_peers_agree_on_version_zero() {
let ours = Hello::v0(16384, 1024);
let theirs = Hello::v0(4096, 512);
let agreed = negotiate(&ours, &theirs).unwrap();
assert_eq!(agreed.version, 0);
assert_eq!(agreed.send_max_header_bytes, 4096);
}
#[test]
fn the_highest_common_version_wins() {
let ours = hello(&[0, 1, 2, 5], &[], &[], 16384);
let theirs = hello(&[1, 2, 3], &[], &[], 16384);
assert_eq!(negotiate(&ours, &theirs).unwrap().version, 2);
}
#[test]
fn version_order_in_the_list_does_not_matter() {
let ours = hello(&[5, 0, 2], &[], &[], 16384);
let theirs = hello(&[2, 5], &[], &[], 16384);
assert_eq!(negotiate(&ours, &theirs).unwrap().version, 5);
}
#[test]
fn an_empty_intersection_fails() {
let ours = hello(&[0], &[], &[], 16384);
for theirs in [hello(&[99], &[], &[], 16384), hello(&[], &[], &[], 16384)] {
assert_eq!(
negotiate(&ours, &theirs),
Err(NegotiateError::NoCommonVersion)
);
}
}
#[test]
fn a_required_capability_we_lack_fails() {
let ours = Hello::v0(16384, 1024);
let theirs = hello(&[0], &[7], &[7], 16384);
assert_eq!(
negotiate(&ours, &theirs),
Err(NegotiateError::UnsupportedRequiredCapability(7))
);
}
#[test]
fn a_required_capability_we_have_succeeds() {
let ours = hello(&[0], &[7, 9], &[], 16384);
let theirs = hello(&[0], &[7], &[7], 8192);
assert_eq!(
negotiate(&ours, &theirs).unwrap(),
Agreed {
version: 0,
send_max_header_bytes: 8192,
guarantees: GuaranteeSet::CORE,
}
);
}
#[test]
fn two_v0_peers_negotiate_core_unchanged() {
let agreed = negotiate(&Hello::v0(16384, 1024), &Hello::v0(16384, 1024)).unwrap();
assert_eq!(agreed.guarantees, GuaranteeSet::CORE);
assert!(agreed.guarantees.is_core());
}
#[test]
fn the_effective_set_is_the_weaker_offer_per_dimension() {
let strong = GuaranteeSet {
ordering: OrderingMode::PerProducerReassemble,
deduplication: Deduplication::Bounded,
dedup_window_ms: Some(60_000),
control_isolated: true,
..GuaranteeSet::CORE
};
let weaker = GuaranteeSet {
ordering: OrderingMode::PerProducerDetect,
deduplication: Deduplication::Bounded,
dedup_window_ms: Some(5_000),
control_isolated: false,
..GuaranteeSet::CORE
};
let agreed = negotiate(
&declaring(strong, GuaranteeSet::CORE),
&declaring(weaker, GuaranteeSet::CORE),
)
.unwrap();
assert_eq!(agreed.guarantees.ordering, OrderingMode::PerProducerDetect);
assert_eq!(agreed.guarantees.dedup_window_ms, Some(5_000));
assert!(!agreed.guarantees.control_isolated);
}
#[test]
fn a_requested_level_the_peer_does_not_offer_fails() {
let wants_ordering = GuaranteeSet {
ordering: OrderingMode::PerProducerDetect,
..GuaranteeSet::CORE
};
let err = negotiate(
&declaring(wants_ordering, wants_ordering),
&Hello::v0(16384, 1024),
)
.unwrap_err();
assert_eq!(
err,
NegotiateError::GuaranteeNotOffered {
peers_requirement: false
}
);
let err = negotiate(
&Hello::v0(16384, 1024),
&declaring(wants_ordering, wants_ordering),
)
.unwrap_err();
assert_eq!(
err,
NegotiateError::GuaranteeNotOffered {
peers_requirement: true
}
);
}
#[test]
fn unordered_dimensions_must_match_exactly() {
for (ours, theirs, dimension) in [
(
GuaranteeSet {
backpressure: Backpressure::Drop,
..GuaranteeSet::CORE
},
GuaranteeSet::CORE,
"backpressure",
),
(
GuaranteeSet {
producer_naming: ProducerNaming::Stable,
..GuaranteeSet::CORE
},
GuaranteeSet::CORE,
"producer naming",
),
(
GuaranteeSet {
acknowledgement: Acknowledgement::Stored,
durability: Some(Durability::Flushed),
..GuaranteeSet::CORE
},
GuaranteeSet {
acknowledgement: Acknowledgement::Stored,
durability: Some(Durability::Written),
..GuaranteeSet::CORE
},
"durability",
),
] {
let err = negotiate(
&declaring(ours, GuaranteeSet::CORE),
&declaring(theirs, GuaranteeSet::CORE),
)
.unwrap_err();
assert_eq!(err, NegotiateError::IncomparableGuarantee { dimension });
}
}
#[test]
fn a_weakened_acknowledgement_drops_the_axes_it_cannot_carry() {
let stored = GuaranteeSet {
acknowledgement: Acknowledgement::Stored,
durability: Some(Durability::Flushed),
..GuaranteeSet::CORE
};
let agreed = negotiate(
&declaring(stored, GuaranteeSet::CORE),
&Hello::v0(16384, 1024),
)
.unwrap();
assert_eq!(
agreed.guarantees.acknowledgement,
Acknowledgement::TransportReceipt
);
assert_eq!(agreed.guarantees.durability, None);
assert!(agreed.guarantees.is_core());
}
#[test]
fn optional_peer_capabilities_are_ignored() {
let ours = Hello::v0(16384, 1024);
let theirs = hello(&[0], &[1, 2, 3], &[], 16384);
assert!(negotiate(&ours, &theirs).is_ok());
}
#[test]
fn version_mismatch_is_checked_before_capabilities() {
let ours = Hello::v0(16384, 1024);
let theirs = hello(&[99], &[], &[42], 16384);
assert_eq!(
negotiate(&ours, &theirs),
Err(NegotiateError::NoCommonVersion)
);
}
#[test]
fn negotiation_is_symmetric_for_the_agreed_version() {
let a = hello(&[0, 1], &[], &[], 1000);
let b = hello(&[1, 2], &[], &[], 2000);
let ab = negotiate(&a, &b).unwrap();
let ba = negotiate(&b, &a).unwrap();
assert_eq!(ab.version, ba.version);
assert_eq!(ab.send_max_header_bytes, 2000);
assert_eq!(ba.send_max_header_bytes, 1000);
}
#[test]
fn errors_become_negotiation_errors() {
let e: weida_core::Error = NegotiateError::NoCommonVersion.into();
assert!(matches!(e, weida_core::Error::Negotiation(_)));
}
}