Skip to main content

fedimint_portalloc/data/
dto.rs

1use std::collections::BTreeMap;
2use std::net::{TcpListener, UdpSocket};
3
4use anyhow::{Result, anyhow};
5use fedimint_core::util::FmtCompact as _;
6use serde::{Deserialize, Serialize};
7use tracing::{debug, trace, warn};
8
9#[cfg(test)]
10mod tests;
11
12/// The lowest port number to try. Ports below 10k are typically used by normal
13/// software, increasing chance they would get in a way.
14const LOW: u16 = 10000;
15
16// The highest port number to try. Ports above 32k are typically ephmeral,
17// increasing a chance of random conflicts after port was already tried.
18const HIGH: u16 = 32000;
19
20const LOG_PORT_ALLOC: &str = "port-alloc";
21
22#[derive(Serialize, Deserialize, Debug, Clone, Default)]
23#[serde(rename_all = "kebab-case")]
24struct RangeData {
25    /// Port range size.
26    size: u16,
27
28    /// Unix timestamp when this range expires.
29    expires: UnixTimestamp,
30}
31
32type UnixTimestamp = u64;
33
34fn default_next() -> u16 {
35    LOW
36}
37
38#[derive(Serialize, Deserialize, Debug, Clone)]
39#[serde(rename_all = "kebab-case")]
40pub struct RootData {
41    /// Next port to try.
42    #[serde(default = "default_next")]
43    next: u16,
44
45    /// Map of port ranges. For each range, the key is the first port in the
46    /// range and the range size and expiration time are stored in the value.
47    keys: BTreeMap<u16, RangeData>,
48}
49
50impl Default for RootData {
51    fn default() -> Self {
52        Self {
53            next: LOW,
54            keys: Default::default(),
55        }
56    }
57}
58
59impl RootData {
60    pub fn get_free_port_range(&mut self, range_size: u16) -> Result<u16> {
61        trace!(target: LOG_PORT_ALLOC, range_size, "Looking for port");
62
63        self.reclaim();
64
65        let mut base_port: u16 = self.next;
66        'retry: loop {
67            trace!(target: LOG_PORT_ALLOC, base_port, range_size, "Checking a port");
68            if base_port > HIGH {
69                self.reclaim();
70                base_port = LOW;
71            }
72            let range = port_range(base_port, range_size)?;
73            if let Some(next_port) = self.contains(range.clone()) {
74                warn!(
75                    base_port,
76                    range_size,
77                    "Could not use a port (already reserved). Will try a different range."
78                );
79                base_port = next_port;
80                continue 'retry;
81            }
82
83            for port in range.clone() {
84                match (
85                    TcpListener::bind(("127.0.0.1", port)),
86                    UdpSocket::bind(("127.0.0.1", port)),
87                ) {
88                    (Err(err), _) | (_, Err(err)) => {
89                        warn!(
90                            err = %err.fmt_compact(),
91                            port, "Could not use a port. Will try a different range"
92                        );
93                        base_port = port + 1;
94                        continue 'retry;
95                    }
96                    (Ok(tcp), Ok(udp)) => (tcp, udp),
97                };
98            }
99
100            self.insert(range);
101            debug!(target: LOG_PORT_ALLOC, base_port, range_size, "Allocated port range");
102            return Ok(base_port);
103        }
104    }
105
106    /// Remove expired entries from the map.
107    fn reclaim(&mut self) {
108        let now = Self::now_ts();
109        self.keys.retain(|_k, v| now < v.expires);
110    }
111
112    /// Check if `range` conflicts with anything already reserved
113    ///
114    /// If it does return next address after the range that conflicted.
115    fn contains(&self, range: std::ops::Range<u16>) -> Option<u16> {
116        self.keys.range(..range.end).next_back().and_then(|(k, v)| {
117            let start = *k;
118            let end = start + v.size;
119
120            if start < range.end && range.start < end {
121                Some(end)
122            } else {
123                None
124            }
125        })
126    }
127
128    fn insert(&mut self, range: std::ops::Range<u16>) {
129        const ALLOCATION_TIME_SECS: u64 = 120;
130
131        // The caller gets some time actually start using the port (`bind`),
132        // to prevent other callers from re-using it. This could typically be
133        // much shorter, as portalloc will not only respect the allocation,
134        // but also try to bind before using a given port range. But for tests
135        // that temporarily release ports (e.g. restarts, failure simulations, etc.),
136        // there's a chance that this can expire and another tests snatches the test,
137        // so better to keep it around the time a longest test can take.
138
139        assert!(self.contains(range.clone()).is_none());
140        self.keys.insert(
141            range.start,
142            RangeData {
143                size: range.len() as u16,
144                expires: Self::now_ts() + ALLOCATION_TIME_SECS,
145            },
146        );
147        self.next = range.end;
148    }
149
150    fn now_ts() -> UnixTimestamp {
151        fedimint_core::time::duration_since_epoch().as_secs()
152    }
153}
154
155fn port_range(base_port: u16, range_size: u16) -> Result<std::ops::Range<u16>> {
156    let end = base_port.checked_add(range_size).ok_or_else(|| {
157        anyhow!("Port range starting at {base_port} with size {range_size} exceeds u16 bounds")
158    })?;
159    Ok(base_port..end)
160}