fedimint_portalloc/data/
dto.rs1use 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
12const LOW: u16 = 10000;
15
16const 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 size: u16,
27
28 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 #[serde(default = "default_next")]
43 next: u16,
44
45 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 fn reclaim(&mut self) {
108 let now = Self::now_ts();
109 self.keys.retain(|_k, v| now < v.expires);
110 }
111
112 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 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}