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