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