1use std::fmt;
9
10use crate::header::{GuaranteeSet, Hello};
11
12#[derive(Clone, Copy, Debug, PartialEq, Eq)]
14pub struct Agreed {
15 pub version: u64,
17 pub send_max_header_bytes: u64,
20 pub guarantees: GuaranteeSet,
26}
27
28#[derive(Clone, Copy, Debug, PartialEq, Eq)]
30pub enum NegotiateError {
31 NoCommonVersion,
33 UnsupportedRequiredCapability(u64),
35 IncomparableGuarantee {
38 dimension: &'static str,
40 },
41 GuaranteeNotOffered {
45 peers_requirement: bool,
47 },
48}
49
50impl fmt::Display for NegotiateError {
51 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
52 match self {
53 NegotiateError::NoCommonVersion => {
54 f.write_str("no wire protocol version is supported by both peers")
55 }
56 NegotiateError::UnsupportedRequiredCapability(c) => {
57 write!(f, "peer requires unsupported capability {c}")
58 }
59 NegotiateError::IncomparableGuarantee { dimension } => {
60 write!(
61 f,
62 "the two peers declare different {dimension}, which is not an ordered dimension"
63 )
64 }
65 NegotiateError::GuaranteeNotOffered { peers_requirement } => {
66 let who = if *peers_requirement { "peer" } else { "local" };
67 write!(
68 f,
69 "the effective guarantee set does not reach the {who} requirement"
70 )
71 }
72 }
73 }
74}
75
76impl std::error::Error for NegotiateError {}
77
78impl From<NegotiateError> for weida_core::Error {
79 fn from(e: NegotiateError) -> weida_core::Error {
80 weida_core::Error::Negotiation(e.to_string())
81 }
82}
83
84pub fn negotiate(ours: &Hello, theirs: &Hello) -> Result<Agreed, NegotiateError> {
95 let version = theirs
96 .versions
97 .iter()
98 .copied()
99 .filter(|v| ours.versions.contains(v))
100 .max()
101 .ok_or(NegotiateError::NoCommonVersion)?;
102
103 if let Some(missing) = theirs
104 .required_capabilities
105 .iter()
106 .copied()
107 .find(|c| !ours.capabilities.contains(c))
108 {
109 return Err(NegotiateError::UnsupportedRequiredCapability(missing));
110 }
111
112 let guarantees = ours
113 .offered()
114 .intersect(&theirs.offered())
115 .map_err(|dimension| NegotiateError::IncomparableGuarantee { dimension })?;
116 for (required, peers_requirement) in [(theirs.required(), true), (ours.required(), false)] {
117 if !guarantees.reaches(&required) {
118 return Err(NegotiateError::GuaranteeNotOffered { peers_requirement });
119 }
120 }
121
122 Ok(Agreed {
123 version,
124 send_max_header_bytes: theirs.max_header_bytes,
125 guarantees,
126 })
127}
128
129#[cfg(test)]
130mod tests {
131 use super::*;
132
133 use crate::header::{
134 Acknowledgement, Backpressure, Deduplication, Durability, OrderingMode, ProducerNaming,
135 };
136
137 fn hello(versions: &[u64], caps: &[u64], required: &[u64], max_header: u64) -> Hello {
138 Hello {
139 versions: versions.to_vec(),
140 max_header_bytes: max_header,
141 max_transfers: 1024,
142 capabilities: caps.to_vec(),
143 required_capabilities: required.to_vec(),
144 guarantees_offered: None,
145 guarantees_required: None,
146 }
147 }
148
149 fn declaring(offered: GuaranteeSet, required: GuaranteeSet) -> Hello {
151 Hello {
152 guarantees_offered: Some(offered),
153 guarantees_required: Some(required),
154 ..Hello::v0(16384, 1024)
155 }
156 }
157
158 #[test]
159 fn v0_peers_agree_on_version_zero() {
160 let ours = Hello::v0(16384, 1024);
161 let theirs = Hello::v0(4096, 512);
162 let agreed = negotiate(&ours, &theirs).unwrap();
163 assert_eq!(agreed.version, 0);
164 assert_eq!(agreed.send_max_header_bytes, 4096);
166 }
167
168 #[test]
169 fn the_highest_common_version_wins() {
170 let ours = hello(&[0, 1, 2, 5], &[], &[], 16384);
171 let theirs = hello(&[1, 2, 3], &[], &[], 16384);
172 assert_eq!(negotiate(&ours, &theirs).unwrap().version, 2);
173 }
174
175 #[test]
176 fn version_order_in_the_list_does_not_matter() {
177 let ours = hello(&[5, 0, 2], &[], &[], 16384);
178 let theirs = hello(&[2, 5], &[], &[], 16384);
179 assert_eq!(negotiate(&ours, &theirs).unwrap().version, 5);
180 }
181
182 #[test]
183 fn an_empty_intersection_fails() {
184 let ours = hello(&[0], &[], &[], 16384);
185 for theirs in [hello(&[99], &[], &[], 16384), hello(&[], &[], &[], 16384)] {
186 assert_eq!(
187 negotiate(&ours, &theirs),
188 Err(NegotiateError::NoCommonVersion)
189 );
190 }
191 }
192
193 #[test]
194 fn a_required_capability_we_lack_fails() {
195 let ours = Hello::v0(16384, 1024);
196 let theirs = hello(&[0], &[7], &[7], 16384);
197 assert_eq!(
198 negotiate(&ours, &theirs),
199 Err(NegotiateError::UnsupportedRequiredCapability(7))
200 );
201 }
202
203 #[test]
204 fn a_required_capability_we_have_succeeds() {
205 let ours = hello(&[0], &[7, 9], &[], 16384);
206 let theirs = hello(&[0], &[7], &[7], 8192);
207 assert_eq!(
208 negotiate(&ours, &theirs).unwrap(),
209 Agreed {
210 version: 0,
211 send_max_header_bytes: 8192,
212 guarantees: GuaranteeSet::CORE,
213 }
214 );
215 }
216
217 #[test]
218 fn two_v0_peers_negotiate_core_unchanged() {
219 let agreed = negotiate(&Hello::v0(16384, 1024), &Hello::v0(16384, 1024)).unwrap();
222 assert_eq!(agreed.guarantees, GuaranteeSet::CORE);
223 assert!(agreed.guarantees.is_core());
224 }
225
226 #[test]
227 fn the_effective_set_is_the_weaker_offer_per_dimension() {
228 let strong = GuaranteeSet {
229 ordering: OrderingMode::PerProducerReassemble,
230 deduplication: Deduplication::Bounded,
231 dedup_window_ms: Some(60_000),
232 control_isolated: true,
233 ..GuaranteeSet::CORE
234 };
235 let weaker = GuaranteeSet {
236 ordering: OrderingMode::PerProducerDetect,
237 deduplication: Deduplication::Bounded,
238 dedup_window_ms: Some(5_000),
239 control_isolated: false,
240 ..GuaranteeSet::CORE
241 };
242 let agreed = negotiate(
243 &declaring(strong, GuaranteeSet::CORE),
244 &declaring(weaker, GuaranteeSet::CORE),
245 )
246 .unwrap();
247 assert_eq!(agreed.guarantees.ordering, OrderingMode::PerProducerDetect);
248 assert_eq!(agreed.guarantees.dedup_window_ms, Some(5_000));
250 assert!(!agreed.guarantees.control_isolated);
251 }
252
253 #[test]
254 fn a_requested_level_the_peer_does_not_offer_fails() {
255 let wants_ordering = GuaranteeSet {
256 ordering: OrderingMode::PerProducerDetect,
257 ..GuaranteeSet::CORE
258 };
259 let err = negotiate(
261 &declaring(wants_ordering, wants_ordering),
262 &Hello::v0(16384, 1024),
263 )
264 .unwrap_err();
265 assert_eq!(
266 err,
267 NegotiateError::GuaranteeNotOffered {
268 peers_requirement: false
269 }
270 );
271 let err = negotiate(
273 &Hello::v0(16384, 1024),
274 &declaring(wants_ordering, wants_ordering),
275 )
276 .unwrap_err();
277 assert_eq!(
278 err,
279 NegotiateError::GuaranteeNotOffered {
280 peers_requirement: true
281 }
282 );
283 }
284
285 #[test]
286 fn unordered_dimensions_must_match_exactly() {
287 for (ours, theirs, dimension) in [
288 (
289 GuaranteeSet {
290 backpressure: Backpressure::Drop,
291 ..GuaranteeSet::CORE
292 },
293 GuaranteeSet::CORE,
294 "backpressure",
295 ),
296 (
297 GuaranteeSet {
298 producer_naming: ProducerNaming::Stable,
299 ..GuaranteeSet::CORE
300 },
301 GuaranteeSet::CORE,
302 "producer naming",
303 ),
304 (
305 GuaranteeSet {
306 acknowledgement: Acknowledgement::Stored,
307 durability: Some(Durability::Flushed),
308 ..GuaranteeSet::CORE
309 },
310 GuaranteeSet {
311 acknowledgement: Acknowledgement::Stored,
312 durability: Some(Durability::Written),
313 ..GuaranteeSet::CORE
314 },
315 "durability",
316 ),
317 ] {
318 let err = negotiate(
319 &declaring(ours, GuaranteeSet::CORE),
320 &declaring(theirs, GuaranteeSet::CORE),
321 )
322 .unwrap_err();
323 assert_eq!(err, NegotiateError::IncomparableGuarantee { dimension });
324 }
325 }
326
327 #[test]
328 fn a_weakened_acknowledgement_drops_the_axes_it_cannot_carry() {
329 let stored = GuaranteeSet {
330 acknowledgement: Acknowledgement::Stored,
331 durability: Some(Durability::Flushed),
332 ..GuaranteeSet::CORE
333 };
334 let agreed = negotiate(
335 &declaring(stored, GuaranteeSet::CORE),
336 &Hello::v0(16384, 1024),
337 )
338 .unwrap();
339 assert_eq!(
342 agreed.guarantees.acknowledgement,
343 Acknowledgement::TransportReceipt
344 );
345 assert_eq!(agreed.guarantees.durability, None);
346 assert!(agreed.guarantees.is_core());
347 }
348
349 #[test]
350 fn optional_peer_capabilities_are_ignored() {
351 let ours = Hello::v0(16384, 1024);
352 let theirs = hello(&[0], &[1, 2, 3], &[], 16384);
353 assert!(negotiate(&ours, &theirs).is_ok());
354 }
355
356 #[test]
357 fn version_mismatch_is_checked_before_capabilities() {
358 let ours = Hello::v0(16384, 1024);
359 let theirs = hello(&[99], &[], &[42], 16384);
360 assert_eq!(
361 negotiate(&ours, &theirs),
362 Err(NegotiateError::NoCommonVersion)
363 );
364 }
365
366 #[test]
367 fn negotiation_is_symmetric_for_the_agreed_version() {
368 let a = hello(&[0, 1], &[], &[], 1000);
369 let b = hello(&[1, 2], &[], &[], 2000);
370 let ab = negotiate(&a, &b).unwrap();
371 let ba = negotiate(&b, &a).unwrap();
372 assert_eq!(ab.version, ba.version);
373 assert_eq!(ab.send_max_header_bytes, 2000);
374 assert_eq!(ba.send_max_header_bytes, 1000);
375 }
376
377 #[test]
378 fn errors_become_negotiation_errors() {
379 let e: weida_core::Error = NegotiateError::NoCommonVersion.into();
380 assert!(matches!(e, weida_core::Error::Negotiation(_)));
381 }
382}