use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::{Mutex, OnceLock};
use std::time::{Duration, Instant};
use netsock::family::AddressFamilyFlags;
use netsock::protocol::ProtocolFlags;
const CACHE_TTL: Duration = Duration::from_secs(60);
pub const OWNER_SCAN_ATTEMPTS: u32 = 3;
pub const OWNER_SCAN_BACKOFF: Duration = Duration::from_millis(10);
pub struct PeerPidLookup {
pub pid: Option<i32>,
pub micros: u64,
}
#[derive(Clone, Copy)]
struct CacheEntry {
pid: i32,
when: Instant,
}
fn cache() -> &'static Mutex<HashMap<SocketAddr, CacheEntry>> {
static C: OnceLock<Mutex<HashMap<SocketAddr, CacheEntry>>> = OnceLock::new();
C.get_or_init(|| Mutex::new(HashMap::new()))
}
pub fn lookup(candidates: &[i32], peer: SocketAddr) -> PeerPidLookup {
let started = Instant::now();
let pid = cached_lookup(candidates, peer);
let micros = u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX);
PeerPidLookup { pid, micros }
}
pub fn lookup_owner(peer: SocketAddr) -> PeerPidLookup {
let started = Instant::now();
let pid = resolve_with_retry(
OWNER_SCAN_ATTEMPTS,
OWNER_SCAN_BACKOFF,
|| owner_scan(peer),
std::thread::sleep,
);
let micros = u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX);
PeerPidLookup { pid, micros }
}
pub async fn lookup_owner_async(peer: SocketAddr) -> PeerPidLookup {
let started = Instant::now();
let pid =
resolve_with_retry_async(OWNER_SCAN_ATTEMPTS, OWNER_SCAN_BACKOFF, || owner_scan(peer))
.await;
let micros = u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX);
PeerPidLookup { pid, micros }
}
pub fn lookup_owner_once(peer: SocketAddr) -> PeerPidLookup {
let started = Instant::now();
let pid = owner_scan(peer);
let micros = u64::try_from(started.elapsed().as_micros()).unwrap_or(u64::MAX);
PeerPidLookup { pid, micros }
}
fn owner_scan(peer: SocketAddr) -> Option<i32> {
#[cfg(target_os = "linux")]
{
cached_lookup_owner(peer)
}
#[cfg(not(target_os = "linux"))]
{
scan_owner(peer)
}
}
fn resolve_with_retry(
attempts: u32,
backoff: Duration,
mut scan: impl FnMut() -> Option<i32>,
mut sleep: impl FnMut(Duration),
) -> Option<i32> {
for attempt in 1..=attempts {
if let Some(pid) = scan() {
return Some(pid);
}
if attempt < attempts {
sleep(backoff);
}
}
None
}
async fn resolve_with_retry_async(
attempts: u32,
backoff: Duration,
mut scan: impl FnMut() -> Option<i32>,
) -> Option<i32> {
for attempt in 1..=attempts {
if let Some(pid) = scan() {
return Some(pid);
}
if attempt < attempts {
tokio::time::sleep(backoff).await;
}
}
None
}
#[cfg(target_os = "linux")]
fn cached_lookup_owner(peer: SocketAddr) -> Option<i32> {
let candidates = same_uid_pids_for_peer_socket(peer)?;
cached_lookup(&candidates, peer)
}
fn cached_lookup(candidates: &[i32], peer: SocketAddr) -> Option<i32> {
let key = peer;
{
let mut map = cache()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(entry) = map.get(&key).copied() {
if entry.when.elapsed() < CACHE_TTL && candidates.contains(&entry.pid) {
return Some(entry.pid);
}
map.remove(&key);
}
}
let pid = scan(candidates, peer)?;
let mut map = cache()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.insert(
key,
CacheEntry {
pid,
when: Instant::now(),
},
);
Some(pid)
}
fn scan(candidates: &[i32], peer: SocketAddr) -> Option<i32> {
let af = AddressFamilyFlags::IPV4 | AddressFamilyFlags::IPV6;
#[cfg(target_os = "linux")]
{
scan_linux(candidates, peer, af)
}
#[cfg(not(target_os = "linux"))]
{
scan_netsock_attached(candidates, peer, af)
}
}
#[cfg(not(target_os = "linux"))]
fn scan_owner(peer: SocketAddr) -> Option<i32> {
let af = AddressFamilyFlags::IPV4 | AddressFamilyFlags::IPV6;
let sockets = netsock::get_sockets(af, ProtocolFlags::TCP).ok()?;
let peer_port = peer.port();
sockets.into_iter().find_map(|s| {
if s.local_port() != peer_port || !addr_matches(s.local_addr(), peer.ip()) {
return None;
}
s.processes
.into_iter()
.find_map(|process| i32::try_from(process.pid).ok())
})
}
#[cfg(target_os = "linux")]
fn scan_linux(candidates: &[i32], peer: SocketAddr, af: AddressFamilyFlags) -> Option<i32> {
let (inode, _uid) = peer_socket_inode_uid(peer, af)?;
candidates
.iter()
.copied()
.find(|&cand| pid_owns_socket_inode(cand, inode))
}
#[cfg(target_os = "linux")]
fn same_uid_pids_for_peer_socket(peer: SocketAddr) -> Option<Vec<i32>> {
let af = AddressFamilyFlags::IPV4 | AddressFamilyFlags::IPV6;
let (_inode, uid) = peer_socket_inode_uid(peer, af)?;
Some(pids_owned_by_uid(uid))
}
#[cfg(target_os = "linux")]
fn peer_socket_inode_uid(peer: SocketAddr, af: AddressFamilyFlags) -> Option<(u32, u32)> {
let sockets = netsock::iter_sockets_without_processes(af, ProtocolFlags::TCP).ok()?;
let peer_port = peer.port();
sockets.into_iter().find_map(|s| {
let s = s.ok()?;
(s.local_port() == peer_port && addr_matches(s.local_addr(), peer.ip()))
.then_some((s.inode, s.uid))
})
}
#[cfg(target_os = "linux")]
fn pids_owned_by_uid(uid: u32) -> Vec<i32> {
use std::os::unix::fs::MetadataExt;
let Ok(entries) = std::fs::read_dir("/proc") else {
return Vec::new();
};
entries
.flatten()
.filter_map(|entry| {
let pid = entry.file_name().to_str()?.parse::<i32>().ok()?;
let metadata = entry.metadata().ok()?;
(metadata.uid() == uid).then_some(pid)
})
.collect()
}
#[cfg(target_os = "linux")]
fn pid_owns_socket_inode(pid: i32, inode: u32) -> bool {
let Ok(entries) = std::fs::read_dir(format!("/proc/{pid}/fd")) else {
return false;
};
let needle = format!("socket:[{inode}]");
entries.flatten().any(|entry| {
std::fs::read_link(entry.path())
.ok()
.is_some_and(|link| link.to_str() == Some(needle.as_str()))
})
}
#[cfg(not(target_os = "linux"))]
fn scan_netsock_attached(
candidates: &[i32],
peer: SocketAddr,
af: AddressFamilyFlags,
) -> Option<i32> {
let sockets = netsock::get_sockets(af, ProtocolFlags::TCP).ok()?;
let peer_port = peer.port();
for s in sockets {
if s.local_port() != peer_port {
continue;
}
if !addr_matches(s.local_addr(), peer.ip()) {
continue;
}
for &cand in candidates {
if let Ok(pid_u32) = u32::try_from(cand)
&& s.is_owned_by_pid(pid_u32)
{
return Some(cand);
}
}
}
None
}
fn addr_matches(local: std::net::IpAddr, peer: std::net::IpAddr) -> bool {
use std::net::IpAddr;
if local == peer {
return true;
}
match (local, peer) {
(IpAddr::V6(v6), IpAddr::V4(v4)) | (IpAddr::V4(v4), IpAddr::V6(v6)) => {
v6.to_ipv4_mapped() == Some(v4)
}
_ => false,
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, TcpListener, TcpStream};
fn clear_cache_for(peer: SocketAddr) {
let mut map = cache()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.remove(&peer);
}
fn lookup_until_pid(candidates: &[i32], peer: SocketAddr, expected: i32) -> PeerPidLookup {
let deadline = Instant::now() + Duration::from_millis(250);
loop {
let got = lookup(candidates, peer);
if got.pid == Some(expected) || Instant::now() >= deadline {
return got;
}
std::thread::sleep(Duration::from_millis(5));
}
}
#[test]
fn addr_matches_same_family() {
let v4 = IpAddr::V4(Ipv4Addr::LOCALHOST);
let v6 = IpAddr::V6(Ipv6Addr::LOCALHOST);
assert!(addr_matches(v4, v4));
assert!(addr_matches(v6, v6));
}
#[test]
fn addr_matches_v4_mapped_v6_either_direction() {
let v4 = IpAddr::V4(Ipv4Addr::LOCALHOST);
let mapped = IpAddr::V6(Ipv4Addr::LOCALHOST.to_ipv6_mapped());
assert!(addr_matches(mapped, v4));
assert!(addr_matches(v4, mapped));
}
#[test]
fn addr_matches_rejects_native_v6_vs_v4() {
assert!(!addr_matches(
IpAddr::V6(Ipv6Addr::LOCALHOST),
IpAddr::V4(Ipv4Addr::LOCALHOST),
));
}
#[test]
fn lookup_finds_self_pid_for_live_loopback_socket() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let me = std::process::id() as i32;
let got = lookup_until_pid(&[me], peer, me);
assert_eq!(
got.pid,
Some(me),
"expected self pid {me} to own loopback socket {peer}",
);
}
#[test]
fn lookup_owner_finds_self_pid_for_live_loopback_socket() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let me = std::process::id() as i32;
let deadline = Instant::now() + Duration::from_millis(250);
let got = loop {
let got = lookup_owner(peer);
if got.pid == Some(me) || Instant::now() >= deadline {
break got;
}
std::thread::sleep(Duration::from_millis(5));
};
assert_eq!(
got.pid,
Some(me),
"expected owner lookup to find self pid {me} for loopback socket {peer}",
);
}
#[test]
fn second_lookup_hits_cache_after_socket_closes() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let peer = {
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let me = std::process::id() as i32;
assert_eq!(
lookup_until_pid(&[me], peer, me).pid,
Some(me),
"primer scan failed"
);
peer
};
let me = std::process::id() as i32;
let got = lookup(&[me], peer);
assert_eq!(
got.pid,
Some(me),
"expected cache hit to keep returning our pid after socket closed",
);
}
#[test]
fn cached_pid_dropped_from_candidates_is_invalidated() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let me = std::process::id() as i32;
assert_eq!(
lookup_until_pid(&[me], peer, me).pid,
Some(me),
"primer scan failed"
);
let got = lookup(&[1], peer);
assert_eq!(
got.pid, None,
"stale cached pid was returned despite being absent from candidates",
);
}
#[test]
fn distinct_addresses_sharing_a_port_do_not_share_a_cache_entry() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let v4_peer = client.local_addr().unwrap();
let v6_peer = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), v4_peer.port());
clear_cache_for(v4_peer);
clear_cache_for(v6_peer);
let me = std::process::id() as i32;
assert_eq!(
lookup_until_pid(&[me], v4_peer, me).pid,
Some(me),
"primer scan failed",
);
assert_eq!(
lookup(&[me], v6_peer).pid,
None,
"the cache entry for {v4_peer} answered for {v6_peer} — \
the cache is keyed by port rather than by address",
);
assert_eq!(lookup(&[me], v4_peer).pid, Some(me));
clear_cache_for(v4_peer);
}
#[test]
fn cache_holds_one_entry_per_address_not_per_port() {
let v4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 9);
let v6 = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 9);
{
let mut map = cache()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.insert(
v4,
CacheEntry {
pid: 111,
when: Instant::now(),
},
);
map.insert(
v6,
CacheEntry {
pid: 222,
when: Instant::now(),
},
);
assert_eq!(map.get(&v4).map(|e| e.pid), Some(111));
assert_eq!(map.get(&v6).map(|e| e.pid), Some(222));
}
clear_cache_for(v4);
clear_cache_for(v6);
}
#[test]
fn retry_returns_first_hit_without_rescanning_or_sleeping() {
let mut calls = 0;
let mut sleeps: Vec<Duration> = Vec::new();
let got = resolve_with_retry(
3,
Duration::from_millis(10),
|| {
calls += 1;
Some(42)
},
|pause| sleeps.push(pause),
);
assert_eq!(got, Some(42));
assert_eq!(calls, 1, "a first-try hit must not trigger extra scans");
assert!(sleeps.is_empty(), "a first-try hit must never sleep");
}
#[test]
fn retry_recovers_from_transient_scan_misses() {
let mut calls = 0;
let mut sleeps: Vec<Duration> = Vec::new();
let got = resolve_with_retry(
3,
Duration::from_millis(10),
|| {
calls += 1;
(calls == 3).then_some(7)
},
|pause| sleeps.push(pause),
);
assert_eq!(got, Some(7));
assert_eq!(calls, 3, "retry should re-scan until the owner appears");
assert_eq!(
sleeps,
vec![Duration::from_millis(10); 2],
"one backoff pause per miss that has an attempt after it"
);
}
#[test]
fn retry_exhausts_attempts_then_fails_closed() {
let mut calls = 0;
let mut sleeps: Vec<Duration> = Vec::new();
let got = resolve_with_retry(
3,
Duration::from_millis(10),
|| {
calls += 1;
None
},
|pause| sleeps.push(pause),
);
assert_eq!(got, None);
assert_eq!(calls, 3, "exactly `attempts` scans, no more, no fewer");
assert_eq!(sleeps.len(), 2, "no sleep after the final miss");
}
#[tokio::test(start_paused = true)]
async fn async_retry_recovers_from_transient_scan_misses_without_blocking() {
let started = tokio::time::Instant::now();
let mut calls = 0;
let got = resolve_with_retry_async(3, Duration::from_millis(10), || {
calls += 1;
(calls == 3).then_some(7)
})
.await;
assert_eq!(got, Some(7));
assert_eq!(calls, 3);
assert_eq!(
started.elapsed(),
Duration::from_millis(20),
"two misses cost exactly two backoff pauses of tokio time"
);
}
#[tokio::test(start_paused = true)]
async fn async_retry_returns_first_hit_without_sleeping() {
let started = tokio::time::Instant::now();
let mut calls = 0;
let got = resolve_with_retry_async(3, Duration::from_millis(10), || {
calls += 1;
Some(42)
})
.await;
assert_eq!(got, Some(42));
assert_eq!(calls, 1);
assert_eq!(started.elapsed(), Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn async_retry_exhausts_attempts_then_fails_closed() {
let started = tokio::time::Instant::now();
let mut calls = 0;
let got = resolve_with_retry_async(3, Duration::from_millis(10), || {
calls += 1;
None
})
.await;
assert_eq!(got, None);
assert_eq!(calls, 3);
assert_eq!(started.elapsed(), Duration::from_millis(20));
}
#[tokio::test]
async fn async_and_once_variants_find_the_same_live_owner() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let me = std::process::id() as i32;
let deadline = Instant::now() + Duration::from_millis(250);
let mut got = lookup_owner_async(peer).await;
while got.pid != Some(me) && Instant::now() < deadline {
got = lookup_owner_async(peer).await;
}
assert_eq!(got.pid, Some(me), "async owner lookup missed {peer}");
let deadline = Instant::now() + Duration::from_millis(250);
let mut once = lookup_owner_once(peer);
while once.pid != Some(me) && Instant::now() < deadline {
once = lookup_owner_once(peer);
}
assert_eq!(once.pid, Some(me), "single-shot owner lookup missed {peer}");
}
#[test]
fn lookup_returns_none_when_no_candidate_matches() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
clear_cache_for(peer);
let got = lookup(&[1], peer);
assert_eq!(got.pid, None);
}
#[test]
#[ignore = "manual benchmark; opt in with --ignored"]
fn bench_lookup_microseconds() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let server_addr = listener.local_addr().unwrap();
let client = TcpStream::connect(server_addr).unwrap();
let (_server_side, _peer) = listener.accept().unwrap();
let peer = client.local_addr().unwrap();
let me = std::process::id() as i32;
const N: usize = 200;
let pct = |samples: &[u64], q: f64| {
samples[((samples.len() as f64 * q) as usize).min(samples.len() - 1)]
};
let mut cold = Vec::with_capacity(N);
for _ in 0..5 {
clear_cache_for(peer);
let _ = lookup(&[me], peer);
}
for _ in 0..N {
clear_cache_for(peer);
cold.push(lookup(&[me], peer).micros);
}
cold.sort_unstable();
eprintln!(
"peer_pid::lookup COLD µs over {N} iters: p50={} p90={} p99={} p999={} max={}",
pct(&cold, 0.50),
pct(&cold, 0.90),
pct(&cold, 0.99),
pct(&cold, 0.999),
cold.last().copied().unwrap_or(0),
);
clear_cache_for(peer);
let _ = lookup(&[me], peer);
let mut warm = Vec::with_capacity(N);
for _ in 0..N {
warm.push(lookup(&[me], peer).micros);
}
warm.sort_unstable();
eprintln!(
"peer_pid::lookup WARM µs over {N} iters: p50={} p90={} p99={} p999={} max={}",
pct(&warm, 0.50),
pct(&warm, 0.90),
pct(&warm, 0.99),
pct(&warm, 0.999),
warm.last().copied().unwrap_or(0),
);
}
}