use anyhow::{Context, Result};
use rusqlite::{Connection, params};
use saorsa_gossip_types::PeerId;
use std::net::SocketAddr;
use std::path::Path;
use std::sync::{Arc, Mutex};
use std::time::SystemTime;
use tracing::{debug, info, warn};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum NatClass {
Public,
FullCone,
RestrictedCone,
PortRestrictedCone,
Symmetric,
Unknown,
}
impl NatClass {
pub fn as_str(&self) -> &'static str {
match self {
NatClass::Public => "public",
NatClass::FullCone => "full_cone",
NatClass::RestrictedCone => "restricted_cone",
NatClass::PortRestrictedCone => "port_restricted_cone",
NatClass::Symmetric => "symmetric",
NatClass::Unknown => "unknown",
}
}
pub fn parse(s: &str) -> Self {
match s {
"public" => NatClass::Public,
"full_cone" => NatClass::FullCone,
"restricted_cone" => NatClass::RestrictedCone,
"port_restricted_cone" => NatClass::PortRestrictedCone,
"symmetric" => NatClass::Symmetric,
_ => NatClass::Unknown,
}
}
}
#[derive(Debug, Clone)]
pub struct PeerCacheEntry {
pub peer_id: PeerId,
pub addr_hints: Vec<SocketAddr>,
pub nat_class: NatClass,
pub roles: Vec<String>, pub last_success: SystemTime,
pub success_count: u32,
pub failure_count: u32,
}
impl PeerCacheEntry {
pub fn score(&self) -> f64 {
let success_rate = if self.success_count + self.failure_count > 0 {
self.success_count as f64 / (self.success_count + self.failure_count) as f64
} else {
0.5 };
let recency_bonus = match self.last_success.elapsed() {
Ok(elapsed) if elapsed.as_secs() < 300 => 1.0, Ok(elapsed) if elapsed.as_secs() < 3600 => 0.7, Ok(elapsed) if elapsed.as_secs() < 86400 => 0.3, _ => 0.0,
};
let nat_penalty = match self.nat_class {
NatClass::Public => 0.0,
NatClass::FullCone => 0.1,
NatClass::RestrictedCone => 0.2,
NatClass::PortRestrictedCone => 0.3,
NatClass::Symmetric => 0.5,
NatClass::Unknown => 0.2,
};
let role_bonus = if self.roles.contains(&"coordinator".to_string()) {
0.3
} else if self.roles.contains(&"relay".to_string()) {
0.2
} else {
0.0
};
(success_rate + recency_bonus + role_bonus - nat_penalty).max(0.0)
}
}
pub struct PeerCache {
conn: Arc<Mutex<Connection>>,
}
impl PeerCache {
const MAX_CACHE_SIZE: usize = 1000;
pub fn default_cache_path() -> Result<std::path::PathBuf> {
let data_dir = dirs::data_local_dir()
.ok_or_else(|| anyhow::anyhow!("Failed to get system data directory"))?
.join("communitas");
std::fs::create_dir_all(&data_dir)
.context("Failed to create communitas data directory")?;
Ok(data_dir.join("peer_cache.db"))
}
pub async fn load(path: &Path) -> Result<Self> {
let conn = Connection::open(path).context("Failed to open peer cache database")?;
conn.pragma_update(None, "journal_mode", "WAL")
.context("Failed to enable WAL mode")?;
conn.pragma_update(None, "busy_timeout", 5000)
.context("Failed to set busy timeout")?;
conn.execute(
"CREATE TABLE IF NOT EXISTS peers (
peer_id TEXT PRIMARY KEY,
addr_hints TEXT NOT NULL,
nat_class TEXT NOT NULL,
roles TEXT NOT NULL,
last_success INTEGER NOT NULL,
success_count INTEGER NOT NULL,
failure_count INTEGER NOT NULL,
is_bootstrap INTEGER NOT NULL DEFAULT 0
)",
[],
)?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_last_success ON peers(last_success)",
[],
)?;
info!("Loaded peer cache from {:?} (WAL mode enabled)", path);
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
})
}
pub async fn update_success(&mut self, peer_id: PeerId, addr: SocketAddr) -> Result<()> {
let conn = self
.conn
.lock()
.map_err(|e| anyhow::anyhow!("Peer cache mutex poisoned: {}", e))?;
let peer_id_str = hex::encode(peer_id.as_bytes());
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.context("Time went backwards")?
.as_secs() as i64;
let exists: bool = conn
.query_row(
"SELECT 1 FROM peers WHERE peer_id = ?1",
params![peer_id_str],
|_| Ok(true),
)
.unwrap_or(false);
if exists {
conn.execute(
"UPDATE peers SET
addr_hints = ?2,
last_success = ?3,
success_count = success_count + 1
WHERE peer_id = ?1",
params![peer_id_str, addr.to_string(), now],
)?;
debug!("Updated peer {} on success", peer_id_str);
} else {
conn.execute(
"INSERT INTO peers (peer_id, addr_hints, nat_class, roles, last_success, success_count, failure_count)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
params![
peer_id_str,
addr.to_string(),
NatClass::Unknown.as_str(),
"[]", now,
1, 0 ],
)?;
debug!("Added new peer {} to cache", peer_id_str);
}
drop(conn); self.enforce_max_size().await?;
Ok(())
}
pub async fn update_failure(&mut self, peer_id: PeerId) -> Result<()> {
let conn = self
.conn
.lock()
.map_err(|e| anyhow::anyhow!("Peer cache mutex poisoned: {}", e))?;
let peer_id_str = hex::encode(peer_id.as_bytes());
conn.execute(
"UPDATE peers SET failure_count = failure_count + 1 WHERE peer_id = ?1",
params![peer_id_str],
)?;
debug!("Incremented failure count for peer {}", peer_id_str);
Ok(())
}
pub fn get_top_peers(&self, limit: usize) -> Vec<PeerCacheEntry> {
let conn = match self.conn.lock() {
Ok(c) => c,
Err(e) => {
warn!("Peer cache mutex poisoned: {}", e);
return Vec::new();
}
};
let mut stmt = conn
.prepare("SELECT peer_id, addr_hints, nat_class, roles, last_success, success_count, failure_count FROM peers")
.ok();
if stmt.is_none() {
return Vec::new();
}
let rows = match stmt.as_mut().and_then(|s| {
s.query_map([], |row| {
let peer_id_hex: String = row.get(0)?;
let peer_id_bytes = hex::decode(&peer_id_hex)
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
let peer_id_array: [u8; 32] = peer_id_bytes
.as_slice()
.try_into()
.map_err(|_| rusqlite::Error::InvalidQuery)?;
let peer_id = PeerId::new(peer_id_array);
let addr_hints_str: String = row.get(1)?;
let addr_hints: Vec<SocketAddr> = addr_hints_str
.split(',')
.filter_map(|s| s.parse().ok())
.collect();
let nat_class = NatClass::parse(&row.get::<_, String>(2)?);
let roles_str: String = row.get(3)?;
let roles: Vec<String> = serde_json::from_str(&roles_str).unwrap_or_default();
let last_success = SystemTime::UNIX_EPOCH
+ std::time::Duration::from_secs(row.get::<_, i64>(4)? as u64);
let success_count: u32 = row.get(5)?;
let failure_count: u32 = row.get(6)?;
Ok(PeerCacheEntry {
peer_id,
addr_hints,
nat_class,
roles,
last_success,
success_count,
failure_count,
})
})
.ok()
}) {
Some(r) => r,
None => return Vec::new(),
};
let mut entries: Vec<PeerCacheEntry> = rows.filter_map(Result::ok).collect();
entries.sort_by(|a, b| {
b.score()
.partial_cmp(&a.score())
.unwrap_or(std::cmp::Ordering::Equal)
});
entries.into_iter().take(limit).collect()
}
pub async fn prune_failed(&mut self, threshold_ratio: f64) -> Result<()> {
let conn = self
.conn
.lock()
.map_err(|e| anyhow::anyhow!("Peer cache mutex poisoned: {}", e))?;
let count = conn.execute(
"DELETE FROM peers WHERE
(failure_count * 1.0) / (success_count + failure_count) > ?1
AND (success_count + failure_count) > 5",
params![threshold_ratio],
)?;
if count > 0 {
warn!("Pruned {} failed peers from cache", count);
}
Ok(())
}
pub fn len(&self) -> usize {
let conn = match self.conn.lock() {
Ok(c) => c,
Err(e) => {
warn!("Peer cache mutex poisoned: {}", e);
return 0;
}
};
conn.query_row("SELECT COUNT(*) FROM peers", [], |row| row.get(0))
.unwrap_or(0)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub async fn add_bootstrap_node(&mut self, four_words: &str) -> Result<()> {
let peer_id_bytes = blake3::hash(four_words.as_bytes());
let peer_id = PeerId::new(*peer_id_bytes.as_bytes());
let conn = self
.conn
.lock()
.map_err(|e| anyhow::anyhow!("Peer cache mutex poisoned: {}", e))?;
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.context("Time went backwards")?
.as_secs() as i64;
let peer_id_str = hex::encode(peer_id.as_bytes());
let exists: bool = conn
.query_row(
"SELECT 1 FROM peers WHERE peer_id = ?1",
params![peer_id_str],
|_| Ok(true),
)
.unwrap_or(false);
if !exists {
conn.execute(
"INSERT INTO peers (peer_id, addr_hints, nat_class, roles, last_success, success_count, failure_count, is_bootstrap)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
params![
peer_id_str,
four_words, NatClass::Public.as_str(), serde_json::to_string(&vec!["bootstrap".to_string()])?,
now,
100, 0, 1 ],
)?;
info!("Added bootstrap node: {}", four_words);
}
Ok(())
}
pub async fn seed_bootstrap_nodes(&mut self, bootstrap_nodes: &[String]) -> Result<usize> {
let mut seeded = 0;
for node in bootstrap_nodes {
match self.add_bootstrap_node(node).await {
Ok(_) => seeded += 1,
Err(e) => warn!("Failed to seed bootstrap node {}: {}", node, e),
}
}
info!("Seeded {} bootstrap nodes into peer cache", seeded);
Ok(seeded)
}
async fn enforce_max_size(&mut self) -> Result<()> {
let conn = self
.conn
.lock()
.map_err(|e| anyhow::anyhow!("Peer cache mutex poisoned: {}", e))?;
let current_size: usize = conn.query_row(
"SELECT COUNT(*) FROM peers",
[],
|row| row.get(0),
)?;
if current_size <= Self::MAX_CACHE_SIZE {
return Ok(());
}
let to_remove = current_size - Self::MAX_CACHE_SIZE;
let removed = conn.execute(
"DELETE FROM peers WHERE peer_id IN (
SELECT peer_id FROM peers
WHERE is_bootstrap = 0
ORDER BY last_success ASC
LIMIT ?1
)",
params![to_remove],
)?;
if removed > 0 {
info!(
"Evicted {} oldest peers (FIFO) to maintain cache size <= {}",
removed,
Self::MAX_CACHE_SIZE
);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
fn create_test_peer_id(seed: u8) -> PeerId {
let mut bytes = [0u8; 32];
bytes[0] = seed;
PeerId::new(bytes)
}
#[tokio::test]
async fn test_peer_cache_creation() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let cache = PeerCache::load(&cache_path)
.await
.expect("Should create new cache");
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[tokio::test]
async fn test_update_success_new_peer() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
cache.update_success(peer_id, addr).await.expect("update");
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
}
#[tokio::test]
async fn test_update_success_existing_peer() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
let peer_id = create_test_peer_id(1);
let addr1: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
let addr2: SocketAddr = "127.0.0.1:8081".parse().expect("addr");
cache
.update_success(peer_id, addr1)
.await
.expect("update 1");
assert_eq!(cache.len(), 1);
cache
.update_success(peer_id, addr2)
.await
.expect("update 2");
assert_eq!(cache.len(), 1);
let peers = cache.get_top_peers(10);
assert_eq!(peers.len(), 1);
assert_eq!(peers[0].success_count, 2);
}
#[tokio::test]
async fn test_update_failure() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
cache.update_success(peer_id, addr).await.expect("success");
cache.update_failure(peer_id).await.expect("failure 1");
cache.update_failure(peer_id).await.expect("failure 2");
let peers = cache.get_top_peers(10);
assert_eq!(peers.len(), 1);
assert_eq!(peers[0].success_count, 1);
assert_eq!(peers[0].failure_count, 2);
}
#[tokio::test]
async fn test_peer_scoring_success_rate() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
let peer_a = create_test_peer_id(1);
let addr_a: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
for _ in 0..10 {
cache.update_success(peer_a, addr_a).await.expect("success");
}
let peer_b = create_test_peer_id(2);
let addr_b: SocketAddr = "127.0.0.1:8081".parse().expect("addr");
for _ in 0..5 {
cache.update_success(peer_b, addr_b).await.expect("success");
cache.update_failure(peer_b).await.expect("failure");
}
let peers = cache.get_top_peers(10);
assert_eq!(peers.len(), 2);
assert_eq!(peers[0].peer_id, peer_a);
assert!(peers[0].score() > peers[1].score());
}
#[tokio::test]
async fn test_peer_scoring_nat_class() {
let entry_public = PeerCacheEntry {
peer_id: create_test_peer_id(1),
addr_hints: vec!["127.0.0.1:8080".parse().expect("addr")],
nat_class: NatClass::Public,
roles: vec![],
last_success: SystemTime::now(),
success_count: 10,
failure_count: 0,
};
let entry_symmetric = PeerCacheEntry {
peer_id: create_test_peer_id(2),
addr_hints: vec!["127.0.0.1:8081".parse().expect("addr")],
nat_class: NatClass::Symmetric,
roles: vec![],
last_success: SystemTime::now(),
success_count: 10,
failure_count: 0,
};
assert!(entry_public.score() > entry_symmetric.score());
}
#[tokio::test]
async fn test_peer_scoring_roles() {
let entry_coordinator = PeerCacheEntry {
peer_id: create_test_peer_id(1),
addr_hints: vec!["127.0.0.1:8080".parse().expect("addr")],
nat_class: NatClass::Unknown,
roles: vec!["coordinator".to_string()],
last_success: SystemTime::now(),
success_count: 10,
failure_count: 0,
};
let entry_regular = PeerCacheEntry {
peer_id: create_test_peer_id(2),
addr_hints: vec!["127.0.0.1:8081".parse().expect("addr")],
nat_class: NatClass::Unknown,
roles: vec![],
last_success: SystemTime::now(),
success_count: 10,
failure_count: 0,
};
assert!(entry_coordinator.score() > entry_regular.score());
}
#[tokio::test]
async fn test_prune_failed_peers() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
let peer_good = create_test_peer_id(1);
let addr_good: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
for _ in 0..10 {
cache
.update_success(peer_good, addr_good)
.await
.expect("success");
}
cache.update_failure(peer_good).await.expect("failure");
let peer_bad = create_test_peer_id(2);
let addr_bad: SocketAddr = "127.0.0.1:8081".parse().expect("addr");
for _ in 0..2 {
cache
.update_success(peer_bad, addr_bad)
.await
.expect("success");
}
for _ in 0..8 {
cache.update_failure(peer_bad).await.expect("failure");
}
assert_eq!(cache.len(), 2);
cache.prune_failed(0.5).await.expect("prune");
assert_eq!(cache.len(), 1);
let peers = cache.get_top_peers(10);
assert_eq!(peers[0].peer_id, peer_good);
}
#[tokio::test]
async fn test_get_top_peers_limit() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
for i in 0..10 {
let peer_id = create_test_peer_id(i);
let addr: SocketAddr = format!("127.0.0.1:80{:02}", i).parse().expect("addr");
cache.update_success(peer_id, addr).await.expect("success");
}
assert_eq!(cache.len(), 10);
let top_5 = cache.get_top_peers(5);
assert_eq!(top_5.len(), 5);
let top_20 = cache.get_top_peers(20);
assert_eq!(top_20.len(), 10);
}
#[tokio::test]
async fn test_persistence_across_restarts() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("addr");
{
let mut cache = PeerCache::load(&cache_path).await.expect("cache 1");
cache.update_success(peer_id, addr).await.expect("success");
assert_eq!(cache.len(), 1);
}
{
let cache = PeerCache::load(&cache_path).await.expect("cache 2");
assert_eq!(cache.len(), 1);
let peers = cache.get_top_peers(10);
assert_eq!(peers[0].peer_id, peer_id);
assert_eq!(peers[0].success_count, 1);
}
}
#[tokio::test]
async fn test_nat_class_serialization() {
let all_classes = vec![
NatClass::Public,
NatClass::FullCone,
NatClass::RestrictedCone,
NatClass::PortRestrictedCone,
NatClass::Symmetric,
NatClass::Unknown,
];
for nat_class in all_classes {
let s = nat_class.as_str();
let deserialized = NatClass::parse(s);
assert_eq!(nat_class, deserialized);
}
assert_eq!(NatClass::parse("invalid"), NatClass::Unknown);
}
#[tokio::test]
async fn test_empty_cache_operations() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.db");
let mut cache = PeerCache::load(&cache_path).await.expect("cache");
cache.prune_failed(0.5).await.expect("prune");
assert_eq!(cache.len(), 0);
let peers = cache.get_top_peers(10);
assert_eq!(peers.len(), 0);
}
#[tokio::test]
async fn test_peer_entry_score_recency() {
use std::time::Duration;
let entry_recent = PeerCacheEntry {
peer_id: create_test_peer_id(1),
addr_hints: vec!["127.0.0.1:8080".parse().expect("addr")],
nat_class: NatClass::Unknown,
roles: vec![],
last_success: SystemTime::now(),
success_count: 10,
failure_count: 0,
};
let entry_old = PeerCacheEntry {
peer_id: create_test_peer_id(2),
addr_hints: vec!["127.0.0.1:8081".parse().expect("addr")],
nat_class: NatClass::Unknown,
roles: vec![],
last_success: SystemTime::now() - Duration::from_secs(86400),
success_count: 10,
failure_count: 0,
};
assert!(entry_recent.score() > entry_old.score());
}
}