1#![deny(missing_docs)]
34#![cfg_attr(docsrs, feature(doc_cfg))]
36
37#[cfg(target_os = "linux")]
40#[path = "platform/linux.rs"]
41mod linux;
42#[cfg(target_os = "macos")]
43#[path = "platform/macos.rs"]
44mod macos;
45#[cfg(target_os = "windows")]
46#[path = "platform/windows.rs"]
47mod windows;
48
49use std::net::SocketAddr;
50use std::sync::Arc;
51use std::time::{Duration, Instant, SystemTime};
52
53use moka::Expiry;
54use moka::{ops::compute::Op, sync::Cache};
55use tokio::{spawn, task::AbortHandle, time::sleep};
56
57#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
59pub struct ConnectionKey {
60 pub local_addr: SocketAddr,
62 pub remote_addr: SocketAddr,
64}
65
66#[derive(Debug, Clone)]
67pub(crate) struct TrackedConnection {
68 pub first_seen: SystemTime,
69 pub last_seen: SystemTime,
70 pub response_count: u64,
71 pub latest_stats: Option<TcpStats>,
72}
73
74struct ExpireAfterTimeout(Duration);
75impl Expiry<ConnectionKey, TrackedConnection> for ExpireAfterTimeout {
76 fn expire_after_create(
77 &self,
78 _key: &ConnectionKey,
79 _value: &TrackedConnection,
80 _created_at: Instant,
81 ) -> Option<Duration> {
82 Some(self.0)
83 }
84
85 fn expire_after_read(
86 &self,
87 _key: &ConnectionKey,
88 value: &TrackedConnection,
89 _read_at: Instant,
90 _duration_until_expiry: Option<Duration>,
91 _last_modified_at: Instant,
92 ) -> Option<Duration> {
93 Some(
94 self.0
95 .saturating_sub(value.last_seen.elapsed().unwrap_or_default()),
96 )
97 }
98
99 fn expire_after_update(
100 &self,
101 _key: &ConnectionKey,
102 value: &TrackedConnection,
103 _updated_at: Instant,
104 _duration_until_expiry: Option<Duration>,
105 ) -> Option<Duration> {
106 Some(
107 self.0
108 .saturating_sub(value.last_seen.elapsed().unwrap_or_default()),
109 )
110 }
111}
112
113#[derive(Debug, Clone, Copy, Default)]
118#[non_exhaustive]
119pub struct TcpStats {
120 pub rtt_us: u32,
122 pub rtt_var_us: u32,
124 pub lost: Option<u32>,
126 pub retrans: u32,
128 pub total_retrans: u32,
130 pub cwnd: u32,
132 pub delivery_rate: Option<u64>,
134}
135
136#[derive(Debug, Clone)]
142#[non_exhaustive]
143pub struct ConnectionSnapshot {
144 pub connection_type: &'static str,
146 pub local_addr: SocketAddr,
148 pub remote_addr: SocketAddr,
150 pub first_seen: SystemTime,
152 pub last_seen: SystemTime,
154 pub expiry: Option<SystemTime>,
156 pub response_count: u64,
158 pub stats: Option<TcpStats>,
163}
164
165type Conns = Cache<ConnectionKey, TrackedConnection>;
166
167#[derive(Debug)]
171pub struct ConnectionTracker {
172 connections: Conns,
173 timeout: Duration,
174 task_abort: AbortHandle,
175}
176
177impl Drop for ConnectionTracker {
178 fn drop(&mut self) {
179 self.task_abort.abort();
180 }
181}
182
183impl ConnectionTracker {
184 pub fn new(timeout: Duration) -> Arc<Self> {
186 let connections = Cache::builder()
187 .expire_after(ExpireAfterTimeout(timeout))
188 .build();
189
190 let conns = connections.clone();
191 let task_abort = spawn(async move {
192 loop {
193 let _ = update_all(conns.clone());
194 sleep(Duration::from_secs(1)).await;
195 }
196 })
197 .abort_handle();
198
199 Arc::new(Self {
200 connections,
201 timeout,
202 task_abort,
203 })
204 }
205
206 pub fn track(&self, local_addr: SocketAddr, remote_addr: SocketAddr) -> bool {
211 let now = SystemTime::now();
212 let key = ConnectionKey {
213 local_addr,
214 remote_addr,
215 };
216 let mut known = false;
217 self.connections.entry(key).and_compute_with(|entry| {
218 if let Some(entry) = entry {
219 known = true;
220 let mut conn = entry.into_value();
221 conn.last_seen = now;
222 conn.response_count += 1;
223 Op::Put(conn)
224 } else {
225 Op::Put(TrackedConnection {
226 first_seen: now,
227 last_seen: now,
228 response_count: 1,
229 latest_stats: None,
230 })
231 }
232 });
233 known
234 }
235
236 pub fn track_warmup(&self, local_addr: SocketAddr, remote_addr: SocketAddr) {
242 let now = SystemTime::now();
243 let key = ConnectionKey {
244 local_addr,
245 remote_addr,
246 };
247 self.connections.entry(key).and_compute_with(|entry| {
248 if entry.is_some() {
249 Op::Nop
250 } else {
251 Op::Put(TrackedConnection {
252 first_seen: now,
253 last_seen: now,
254 response_count: 0,
255 latest_stats: None,
256 })
257 }
258 });
259 }
260
261 pub fn snapshot(&self) -> Vec<ConnectionSnapshot> {
263 self.connections
264 .iter()
265 .map(|(key, conn)| ConnectionSnapshot {
266 connection_type: "tcp",
267 local_addr: key.local_addr,
268 remote_addr: key.remote_addr,
269 first_seen: conn.first_seen,
270 last_seen: conn.last_seen,
271 expiry: conn.last_seen.checked_add(self.timeout),
272 response_count: conn.response_count,
273 stats: conn.latest_stats,
274 })
275 .collect()
276 }
277}
278
279fn update_all(conns: Conns) -> std::io::Result<()> {
280 let keys: Vec<ConnectionKey> = conns.iter().map(|(k, _)| *k).collect();
281 if keys.is_empty() {
282 return Ok(());
283 }
284
285 #[allow(
286 unused_variables,
287 reason = "when any of the platform-specific impls work, this will be shadowed"
288 )]
289 let stats: Vec<(ConnectionKey, TcpStats)> = Vec::new();
290
291 #[cfg(target_os = "linux")]
292 let stats = linux::query_tcp_stats(&keys)?;
293
294 #[cfg(target_os = "macos")]
295 let stats = macos::query_tcp_stats(&keys)?;
296
297 #[cfg(target_os = "windows")]
298 let stats = windows::query_tcp_stats(&keys)?;
299
300 for (key, tcp_stats) in &stats {
301 update_stats(&conns, *key, *tcp_stats);
302 }
303
304 Ok(())
305}
306
307fn update_stats(conns: &Conns, key: ConnectionKey, stats: TcpStats) {
308 conns.entry(key).and_compute_with(|entry| {
309 if let Some(entry) = entry {
310 let mut entry = entry.into_value();
311 entry.latest_stats = Some(stats);
312 Op::Put(entry)
313 } else {
314 Op::Nop
315 }
316 });
317}