1use std::collections::HashMap;
7use std::net::SocketAddr;
8use std::time::Duration;
9
10use bytes::Bytes;
11use rand::rngs::StdRng;
12use rand::{Rng, SeedableRng};
13use sipx_sip::headers::{OcParameter, OverloadAlgorithm, OverloadSequence, Via, first_hop_end};
14use sipx_sip::{Header, HeaderName, Headers, Request, Response};
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum RequestCategory {
19 Ordinary,
21 Protected,
23}
24
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum OverloadFeedback {
28 Loss(u8),
30 Rate(u32),
32}
33
34#[derive(Debug, Clone)]
36pub struct OverloadConfig {
37 pub advertise: bool,
42 pub feedback: OverloadFeedback,
44 pub validity: Duration,
46 pub rate_tolerance_intervals: u32,
48 pub rate_priority_tolerance_intervals: u32,
50 pub peer_limit: usize,
52 pub categorize: fn(&Request) -> RequestCategory,
54}
55
56impl Default for OverloadConfig {
57 fn default() -> Self {
58 Self {
59 advertise: false,
60 feedback: OverloadFeedback::Loss(100),
61 validity: Duration::from_millis(500),
62 rate_tolerance_intervals: 5,
63 rate_priority_tolerance_intervals: 10,
64 peer_limit: 1024,
65 categorize: |_| RequestCategory::Ordinary,
66 }
67 }
68}
69
70#[derive(Debug, Clone)]
71pub(crate) struct Report {
72 algorithm: OverloadAlgorithm,
73 value: u32,
74 validity: Option<Duration>,
75 sequence: OverloadSequence,
76}
77
78#[derive(Debug)]
79enum Active {
80 Loss {
81 percentage: f64,
82 ordinary: u32,
83 protected: u32,
84 },
85 Rate(RateState),
86}
87
88#[derive(Debug)]
89struct PeerState {
90 sequence: OverloadSequence,
91 until: Option<Duration>,
92 active: Option<Active>,
93 last_used: Duration,
94}
95
96#[derive(Debug)]
97struct RateState {
98 interval: f64,
99 ordinary_tolerance: f64,
100 protected_tolerance: f64,
101 content: f64,
102 last_forwarded: Duration,
103}
104
105impl RateState {
106 fn new(
107 rate: u32,
108 tolerance_intervals: u32,
109 priority_tolerance_intervals: u32,
110 now: Duration,
111 ) -> Self {
112 if rate == 0 {
113 return Self {
114 interval: f64::INFINITY,
115 ordinary_tolerance: 0.0,
116 protected_tolerance: 0.0,
117 content: 0.0,
118 last_forwarded: now,
119 };
120 }
121 let interval = 1.0 / f64::from(rate);
122 Self {
123 interval,
124 ordinary_tolerance: interval * f64::from(tolerance_intervals),
125 protected_tolerance: interval * f64::from(priority_tolerance_intervals),
126 content: 0.0,
127 last_forwarded: now,
128 }
129 }
130
131 fn admit(&mut self, now: Duration, category: RequestCategory) -> bool {
132 if self.interval.is_infinite() {
133 return false;
134 }
135 let elapsed = now.saturating_sub(self.last_forwarded).as_secs_f64();
136 let provisional = (self.content - elapsed).max(0.0);
137 let tolerance = match category {
138 RequestCategory::Ordinary => self.ordinary_tolerance,
139 RequestCategory::Protected => self.protected_tolerance,
140 };
141 if provisional > tolerance {
142 return false;
143 }
144 self.content = provisional + self.interval;
145 self.last_forwarded = now;
146 true
147 }
148}
149
150#[derive(Debug)]
152pub(crate) struct Controller {
153 peers: HashMap<SocketAddr, PeerState>,
154 random: StdRng,
155 tolerance_intervals: u32,
156 priority_tolerance_intervals: u32,
157 peer_limit: usize,
158}
159
160impl Controller {
161 pub(crate) fn new(
162 tolerance_intervals: u32,
163 priority_tolerance_intervals: u32,
164 peer_limit: usize,
165 ) -> Self {
166 Self {
167 peers: HashMap::new(),
168 random: StdRng::from_os_rng(),
169 tolerance_intervals,
170 priority_tolerance_intervals,
171 peer_limit,
172 }
173 }
174
175 #[cfg(test)]
176 fn seeded(seed: u64, tolerance_intervals: u32, priority_tolerance_intervals: u32) -> Self {
177 Self {
178 peers: HashMap::new(),
179 random: StdRng::seed_from_u64(seed),
180 tolerance_intervals,
181 priority_tolerance_intervals,
182 peer_limit: 1024,
183 }
184 }
185
186 #[cfg(test)]
187 fn seeded_with_limit(
188 seed: u64,
189 tolerance_intervals: u32,
190 priority_tolerance_intervals: u32,
191 peer_limit: usize,
192 ) -> Self {
193 Self {
194 peers: HashMap::new(),
195 random: StdRng::seed_from_u64(seed),
196 tolerance_intervals,
197 priority_tolerance_intervals,
198 peer_limit,
199 }
200 }
201
202 pub(crate) fn observe(&mut self, peer: SocketAddr, response: &Response, now: Duration) {
203 if let Some(report) = report_from(response) {
204 self.apply(peer, &report, now);
205 }
206 }
207
208 fn apply(&mut self, peer: SocketAddr, report: &Report, now: Duration) {
209 if self
210 .peers
211 .get(&peer)
212 .is_some_and(|state| report.sequence <= state.sequence)
213 {
214 return;
215 }
216
217 self.make_room_for(peer, now);
218
219 let validity = report.validity.unwrap_or(Duration::from_millis(500));
220 if validity.is_zero() {
221 self.peers.insert(
222 peer,
223 PeerState {
224 sequence: report.sequence,
225 until: None,
226 active: None,
227 last_used: now,
228 },
229 );
230 return;
231 }
232
233 let active = match &report.algorithm {
234 OverloadAlgorithm::Loss if report.value <= 100 => Some(Active::Loss {
235 percentage: f64::from(report.value),
236 ordinary: 0,
237 protected: 0,
238 }),
239 OverloadAlgorithm::Rate => Some(Active::Rate(RateState::new(
240 report.value,
241 self.tolerance_intervals,
242 self.priority_tolerance_intervals,
243 now,
244 ))),
245 OverloadAlgorithm::Loss | OverloadAlgorithm::Other(_) => None,
246 };
247 self.peers.insert(
248 peer,
249 PeerState {
250 sequence: report.sequence,
251 until: Some(now.saturating_add(validity)),
252 active,
253 last_used: now,
254 },
255 );
256 }
257
258 fn make_room_for(&mut self, peer: SocketAddr, now: Duration) {
259 if self.peers.contains_key(&peer) || self.peers.len() < self.peer_limit {
260 return;
261 }
262 let candidate = self
263 .peers
264 .iter()
265 .min_by_key(|(_, state)| {
266 let active = state
267 .active
268 .as_ref()
269 .is_some_and(|_| state.until.is_some_and(|until| now < until));
270 (active, state.last_used)
271 })
272 .map(|(peer, _)| *peer);
273 if let Some(candidate) = candidate {
274 self.peers.remove(&candidate);
275 }
276 }
277
278 pub(crate) fn admit(
279 &mut self,
280 peer: SocketAddr,
281 category: RequestCategory,
282 now: Duration,
283 ) -> bool {
284 let Some(state) = self.peers.get_mut(&peer) else {
285 return true;
286 };
287 state.last_used = now;
288 if state.until.is_none_or(|until| now >= until) {
289 state.active = None;
290 state.until = None;
291 return true;
292 }
293 let Some(active) = state.active.as_mut() else {
294 return true;
295 };
296 match active {
297 Active::Rate(rate) => rate.admit(now, category),
298 Active::Loss {
299 percentage,
300 ordinary,
301 protected,
302 } => {
303 match category {
304 RequestCategory::Ordinary => *ordinary = ordinary.saturating_add(1),
305 RequestCategory::Protected => *protected = protected.saturating_add(1),
306 }
307 let total = f64::from(ordinary.saturating_add(*protected));
308 let ordinary_share = (f64::from(*ordinary) / total) * 100.0;
309 let protected_share = (f64::from(*protected) / total) * 100.0;
310 let discard = match category {
311 RequestCategory::Ordinary if *percentage <= ordinary_share => {
312 *percentage / ordinary_share
313 }
314 RequestCategory::Ordinary => 1.0,
315 RequestCategory::Protected if *percentage <= ordinary_share => 0.0,
316 RequestCategory::Protected if protected_share > 0.0 => {
317 (*percentage - ordinary_share) / protected_share
318 }
319 RequestCategory::Protected => 0.0,
320 };
321 self.random.random::<f64>() >= discard.clamp(0.0, 1.0)
322 }
323 }
324 }
325}
326
327pub(crate) fn advertise(request: &mut Request) {
329 rewrite_top_via(&mut request.headers, b";oc;oc-algo=\"loss,rate\"");
330}
331
332pub(crate) fn add_feedback(
334 response: &mut Response,
335 request: &Request,
336 feedback: OverloadFeedback,
337 validity: Duration,
338 sequence: OverloadSequence,
339) -> bool {
340 let offered = request
341 .headers
342 .typed::<Via>()
343 .and_then(Result::ok)
344 .and_then(|via| via.overload().ok())
345 .is_some_and(|parameters| {
346 parameters.oc == Some(OcParameter::Support)
347 && parameters.algorithms.iter().any(|algorithm| {
348 matches!(
349 (feedback, algorithm),
350 (OverloadFeedback::Loss(_), OverloadAlgorithm::Loss)
351 | (OverloadFeedback::Rate(_), OverloadAlgorithm::Rate)
352 )
353 })
354 });
355 if !offered {
356 return false;
357 }
358 let (value, algorithm) = match feedback {
359 OverloadFeedback::Loss(value) => (u64::from(value), "loss"),
360 OverloadFeedback::Rate(value) => (u64::from(value), "rate"),
361 };
362 let validity = u64::try_from(validity.as_millis()).unwrap_or(u64::MAX);
363 let addition =
364 format!(";oc={value};oc-algo=\"{algorithm}\";oc-validity={validity};oc-seq={sequence}");
365 rewrite_top_via(&mut response.headers, addition.as_bytes())
366}
367
368fn report_from(response: &Response) -> Option<Report> {
369 let parameters = response.headers.typed::<Via>()?.ok()?.overload().ok()?;
370 let OcParameter::Value(value) = parameters.oc? else {
371 return None;
372 };
373 let [algorithm] = parameters.algorithms.as_slice() else {
374 return None;
375 };
376 Some(Report {
377 algorithm: algorithm.clone(),
378 value: u32::try_from(value).ok()?,
379 validity: parameters.validity,
380 sequence: parameters.sequence?,
381 })
382}
383
384fn rewrite_top_via(headers: &mut Headers, addition: &[u8]) -> bool {
385 let Some(header) = headers.get(&HeaderName::Via) else {
386 return false;
387 };
388 let value = header.value().into_owned();
389 let hop_end = first_hop_end(&value);
390 let mut hop = value.get(..hop_end).unwrap_or(&value).to_vec();
391 if Via::parse_one(&hop).is_err() {
392 return false;
393 }
394 for name in [b"oc".as_slice(), b"oc-algo", b"oc-validity", b"oc-seq"] {
395 while let Some((start, end)) = crate::nat::param_span(&hop, name) {
396 hop.drain(start..end);
397 }
398 }
399 hop.extend_from_slice(addition);
400 let mut rebuilt = Vec::with_capacity(value.len().saturating_add(addition.len()));
401 rebuilt.extend_from_slice(&hop);
402 rebuilt.extend_from_slice(value.get(hop_end..).unwrap_or(&[]));
403 let Ok(header) = Header::build(HeaderName::Via, Bytes::from(rebuilt)) else {
404 return false;
405 };
406 if headers.remove_first(&HeaderName::Via).is_none() {
407 return false;
408 }
409 headers.push_front(header);
410 true
411}
412
413#[cfg(test)]
414#[allow(
415 clippy::unwrap_used,
416 clippy::expect_used,
417 clippy::panic,
418 clippy::indexing_slicing
419)]
420mod tests {
421 use std::net::SocketAddr;
422 use std::time::Duration;
423
424 use bytes::Bytes;
425 use sipx_sip::headers::{OcParameter, OverloadAlgorithm, OverloadSequence, Via};
426 use sipx_sip::{Limits, Message, StatusCode, parse_datagram};
427
428 use super::{Controller, OverloadFeedback, Report, RequestCategory, add_feedback, advertise};
429
430 fn peer() -> SocketAddr {
431 "192.0.2.10:5060".parse().expect("peer")
432 }
433
434 fn sequence(value: u64) -> OverloadSequence {
435 OverloadSequence::from_integer(value).expect("small sequence")
436 }
437
438 fn report(algorithm: OverloadAlgorithm, value: u32, sequence_number: u64) -> Report {
439 Report {
440 algorithm,
441 value,
442 validity: Some(Duration::from_secs(10)),
443 sequence: sequence(sequence_number),
444 }
445 }
446
447 #[test]
450 fn a_client_told_to_reduce_by_half_forwards_half_as_many_requests() {
451 let mut controller = Controller::seeded(0x7339, 0, 0);
452 controller.apply(
453 peer(),
454 &report(OverloadAlgorithm::Loss, 50, 1),
455 Duration::ZERO,
456 );
457
458 let forwarded = (0..10_000)
459 .filter(|_| {
460 controller.admit(peer(), RequestCategory::Ordinary, Duration::from_millis(1))
461 })
462 .count();
463 assert!(
464 (4_900..=5_100).contains(&forwarded),
465 "50% loss forwarded {forwarded}/10000 for the fixed seed"
466 );
467 }
468
469 #[test]
470 fn protected_requests_survive_while_ordinary_traffic_can_supply_the_reduction() {
471 let mut controller = Controller::seeded(9, 0, 0);
472 controller.apply(
473 peer(),
474 &report(OverloadAlgorithm::Loss, 50, 1),
475 Duration::ZERO,
476 );
477 for _ in 0..80 {
478 let _ = controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO);
479 }
480 for _ in 0..20 {
481 assert!(controller.admit(peer(), RequestCategory::Protected, Duration::ZERO));
482 }
483 }
484
485 #[test]
486 fn stale_reports_and_expired_reports_do_not_control_the_client() {
487 let mut controller = Controller::seeded(2, 0, 0);
488 controller.apply(
489 peer(),
490 &report(OverloadAlgorithm::Loss, 100, 2),
491 Duration::ZERO,
492 );
493 controller.apply(
494 peer(),
495 &report(OverloadAlgorithm::Loss, 0, 1),
496 Duration::ZERO,
497 );
498 assert!(!controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
499
500 controller.apply(
501 peer(),
502 &Report {
503 algorithm: OverloadAlgorithm::Loss,
504 value: 100,
505 validity: Some(Duration::ZERO),
506 sequence: sequence(3),
507 },
508 Duration::ZERO,
509 );
510 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
511
512 controller.apply(
513 peer(),
514 &report(OverloadAlgorithm::Loss, 100, 2),
515 Duration::ZERO,
516 );
517 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
518
519 controller.apply(
520 peer(),
521 &Report {
522 validity: Some(Duration::from_millis(10)),
523 sequence: sequence(4),
524 ..report(OverloadAlgorithm::Loss, 100, 4)
525 },
526 Duration::ZERO,
527 );
528 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::from_millis(10)));
529 }
530
531 #[test]
532 fn absent_validity_means_five_hundred_milliseconds_and_a_newer_report_restarts_control() {
533 let mut controller = Controller::seeded(4, 0, 0);
534 controller.apply(
535 peer(),
536 &Report {
537 algorithm: OverloadAlgorithm::Loss,
538 value: 100,
539 validity: None,
540 sequence: sequence(1),
541 },
542 Duration::ZERO,
543 );
544 assert!(!controller.admit(
545 peer(),
546 RequestCategory::Ordinary,
547 Duration::from_millis(499)
548 ));
549 assert!(controller.admit(
550 peer(),
551 RequestCategory::Ordinary,
552 Duration::from_millis(500)
553 ));
554
555 controller.apply(
556 peer(),
557 &report(OverloadAlgorithm::Loss, 100, 2),
558 Duration::from_millis(501),
559 );
560 assert!(!controller.admit(
561 peer(),
562 RequestCategory::Ordinary,
563 Duration::from_millis(501)
564 ));
565 }
566
567 #[test]
568 fn rate_control_paces_against_supplied_time() {
569 let mut controller = Controller::seeded(3, 0, 0);
570 controller.apply(
571 peer(),
572 &report(OverloadAlgorithm::Rate, 2, 1),
573 Duration::ZERO,
574 );
575 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
576 assert!(!controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
577 assert!(!controller.admit(
578 peer(),
579 RequestCategory::Ordinary,
580 Duration::from_millis(499)
581 ));
582 assert!(controller.admit(
583 peer(),
584 RequestCategory::Ordinary,
585 Duration::from_millis(500)
586 ));
587 }
588
589 #[test]
590 fn rate_burst_tolerance_is_an_input_not_a_hidden_constant() {
591 let mut controller = Controller::seeded(5, 2, 2);
592 controller.apply(
593 peer(),
594 &report(OverloadAlgorithm::Rate, 2, 1),
595 Duration::ZERO,
596 );
597 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
598 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
599 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
600 assert!(!controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
601 }
602
603 #[test]
604 fn rate_priority_uses_the_second_threshold_after_ordinary_traffic_is_blocked() {
605 let mut controller = Controller::seeded(7, 0, 2);
606 controller.apply(
607 peer(),
608 &report(OverloadAlgorithm::Rate, 2, 1),
609 Duration::ZERO,
610 );
611 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
612 assert!(!controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
613 assert!(controller.admit(peer(), RequestCategory::Protected, Duration::ZERO));
614 }
615
616 #[test]
617 fn zero_rate_rejects_everything_but_zero_validity_disables_control() {
618 let mut controller = Controller::seeded(6, 0, 0);
619 controller.apply(
620 peer(),
621 &report(OverloadAlgorithm::Rate, 0, 1),
622 Duration::ZERO,
623 );
624 assert!(!controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
625 controller.apply(
626 peer(),
627 &Report {
628 algorithm: OverloadAlgorithm::Rate,
629 value: 0,
630 validity: Some(Duration::ZERO),
631 sequence: sequence(2),
632 },
633 Duration::ZERO,
634 );
635 assert!(controller.admit(peer(), RequestCategory::Ordinary, Duration::ZERO));
636 }
637
638 #[test]
639 fn feedback_from_many_peers_never_exceeds_the_configured_state_bound() {
640 let mut controller = Controller::seeded_with_limit(8, 0, 0, 4);
641
642 for port in 5000..5100 {
643 let peer = SocketAddr::from(([192, 0, 2, 10], port));
644 controller.apply(
645 peer,
646 &report(OverloadAlgorithm::Loss, 100, u64::from(port)),
647 Duration::from_millis(u64::from(port)),
648 );
649 assert!(
650 controller.peers.len() <= 4,
651 "peer state exceeded its configured bound"
652 );
653 }
654
655 assert_eq!(controller.peers.len(), 4);
656 }
657
658 #[test]
659 fn client_and_server_generate_only_the_parameters_their_role_owns() {
660 let bytes = Bytes::from_static(
661 b"OPTIONS sip:a@example SIP/2.0\r\n\
662 Via: SIP/2.0/UDP client.example;branch=z9hG4bKx;oc=9;oc-algo=loss;\
663 oc-validity=1;oc-seq=1.0\r\n\
664 To: <sip:a@example>\r\n\
665 From: <sip:b@example>;tag=one\r\n\
666 Call-ID: roles@sipx\r\n\
667 CSeq: 1 OPTIONS\r\n\
668 Max-Forwards: 70\r\n\
669 Content-Length: 0\r\n\r\n",
670 );
671 let Message::Request(mut request) =
672 parse_datagram(bytes, &Limits::datagram()).expect("request parses")
673 else {
674 panic!("request expected");
675 };
676 advertise(&mut request);
677 let offer = request
678 .headers
679 .typed::<Via>()
680 .expect("Via")
681 .expect("Via parses")
682 .overload()
683 .expect("offer parses");
684 assert_eq!(offer.oc, Some(OcParameter::Support));
685 assert_eq!(
686 offer.algorithms,
687 vec![OverloadAlgorithm::Loss, OverloadAlgorithm::Rate]
688 );
689 assert_eq!(offer.validity, None);
690 assert_eq!(offer.sequence, None);
691
692 let status = StatusCode::new(503).expect("status");
693 let mut response =
694 sipx_sip::ResponseBuilder::to_request(&request, status, "Service Unavailable")
695 .expect("response")
696 .build();
697 assert!(add_feedback(
698 &mut response,
699 &request,
700 OverloadFeedback::Rate(150),
701 Duration::from_secs(1),
702 sequence(2),
703 ));
704 let report = response
705 .headers
706 .typed::<Via>()
707 .expect("Via")
708 .expect("Via parses")
709 .overload()
710 .expect("report parses");
711 assert_eq!(report.oc, Some(OcParameter::Value(150)));
712 assert_eq!(report.algorithms, vec![OverloadAlgorithm::Rate]);
713 assert_eq!(report.validity, Some(Duration::from_secs(1)));
714 assert_eq!(report.sequence, Some(sequence(2)));
715 assert!(
716 response
717 .headers
718 .value(&sipx_sip::HeaderName::Via)
719 .is_some_and(|value| {
720 value
721 .windows(b"oc-algo=\"rate\"".len())
722 .any(|part| part == b"oc-algo=\"rate\"")
723 }),
724 "the server algorithm token is a quoted algo-list on the wire"
725 );
726 }
727}