Skip to main content

sipx_transport/
overload.rs

1//! Hop-by-hop overload control (RFC 7339 and RFC 7415).
2//!
3//! The arithmetic takes elapsed time and seeded randomness as inputs. The endpoint driver owns the
4//! controller and supplies both, keeping response updates and request admission on its serial loop.
5
6use 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/// Which RFC 7339 message category local policy assigns to a request.
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum RequestCategory {
19    /// Traffic reduced first.
20    Ordinary,
21    /// In-dialog, emergency, or other locally important traffic reduced only when necessary.
22    Protected,
23}
24
25/// The feedback an endpoint reports when its application queue is full.
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum OverloadFeedback {
28    /// Ask the upstream peer to discard this percentage of requests.
29    Loss(u8),
30    /// Cap the upstream peer at this many requests per second.
31    Rate(u32),
32}
33
34/// Endpoint overload-control policy.
35#[derive(Debug, Clone)]
36pub struct OverloadConfig {
37    /// Advertise RFC 7339/7415 client support on every outgoing topmost `Via`.
38    ///
39    /// Off by default because overload control is an extension negotiated hop by hop. Enabling it
40    /// retains the complete `loss,rate` offer; it does not select a compatibility subset per peer.
41    pub advertise: bool,
42    /// What to report upstream when the application queue is full.
43    pub feedback: OverloadFeedback,
44    /// How long that report remains valid.
45    pub validity: Duration,
46    /// RFC 7415's `TAU1` for ordinary requests, in target inter-request intervals.
47    pub rate_tolerance_intervals: u32,
48    /// RFC 7415's larger `TAU2` for protected requests, in target intervals.
49    pub rate_priority_tolerance_intervals: u32,
50    /// Most downstream peers whose feedback sequence and algorithm state are retained.
51    pub peer_limit: usize,
52    /// Assign a request to one of RFC 7339 ยง7.2's two categories.
53    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/// Per-next-hop client state. Owned and called only by the endpoint driver.
151#[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
327/// Add the client capability parameters and remove server-only parameters from the top Via.
328pub(crate) fn advertise(request: &mut Request) {
329    rewrite_top_via(&mut request.headers, b";oc;oc-algo=\"loss,rate\"");
330}
331
332/// Add server feedback if the request offered the configured algorithm.
333pub(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    /// T-22's failing-first witness. A fixed seed makes the actual reduction reviewable instead of
448    /// accepting whichever distribution an operating-system generator happened to produce.
449    #[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}