use anyhow::{Context, Result};
use saorsa_gossip_types::PeerId;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::path::{Path, PathBuf};
use std::time::SystemTime;
use tracing::{debug, info, warn};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
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, Serialize, Deserialize)]
pub struct PeerCacheEntry {
#[serde(
serialize_with = "serialize_peer_id",
deserialize_with = "deserialize_peer_id"
)]
pub peer_id: PeerId,
pub addr_hints: Vec<SocketAddr>,
pub nat_class: NatClass,
pub roles: Vec<String>,
#[serde(
serialize_with = "serialize_time",
deserialize_with = "deserialize_time"
)]
pub last_success: SystemTime,
pub success_count: u32,
pub failure_count: u32,
#[serde(default)]
pub is_bootstrap: bool,
}
fn serialize_peer_id<S>(peer_id: &PeerId, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str(&hex::encode(peer_id.as_bytes()))
}
fn deserialize_peer_id<'de, D>(deserializer: D) -> Result<PeerId, D::Error>
where
D: serde::Deserializer<'de>,
{
let hex_str = String::deserialize(deserializer)?;
let bytes = hex::decode(&hex_str).map_err(serde::de::Error::custom)?;
let array: [u8; 32] = bytes
.as_slice()
.try_into()
.map_err(|_| serde::de::Error::custom("Invalid PeerId length"))?;
Ok(PeerId::new(array))
}
fn serialize_time<S>(time: &SystemTime, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
let secs = time
.duration_since(SystemTime::UNIX_EPOCH)
.map_err(serde::ser::Error::custom)?
.as_secs();
serializer.serialize_u64(secs)
}
fn deserialize_time<'de, D>(deserializer: D) -> Result<SystemTime, D::Error>
where
D: serde::Deserializer<'de>,
{
let secs = u64::deserialize(deserializer)?;
Ok(SystemTime::UNIX_EPOCH + std::time::Duration::from_secs(secs))
}
impl PeerCacheEntry {
pub fn score(&self) -> f64 {
let total_attempts = self.success_count + self.failure_count;
if total_attempts == 0 {
return 0.0;
}
let success_rate = self.success_count as f64 / total_attempts as f64;
let now = SystemTime::now();
let age_secs = now
.duration_since(self.last_success)
.map(|d| d.as_secs())
.unwrap_or(u64::MAX);
const ONE_DAY: u64 = 24 * 60 * 60;
let recency_bonus = (-(age_secs as f64 / ONE_DAY as f64)).exp();
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.25, };
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)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct PeerCacheFile {
version: u32,
peers: HashMap<String, PeerCacheEntry>,
}
impl PeerCacheFile {
fn new() -> Self {
Self {
version: 1,
peers: HashMap::new(),
}
}
}
pub struct PeerCache {
cache_path: PathBuf,
}
impl PeerCache {
const MAX_CACHE_SIZE: usize = 1000;
pub fn default_cache_path() -> Result<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.json"))
}
pub async fn load(path: &Path) -> Result<Self> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)
.with_context(|| format!("Failed to create directory: {:?}", parent))?;
}
if !path.exists() {
let empty_cache = PeerCacheFile::new();
let json = serde_json::to_string_pretty(&empty_cache)
.context("Failed to serialize empty cache")?;
std::fs::write(path, json)
.with_context(|| format!("Failed to write cache file: {:?}", path))?;
info!("Created new peer cache at {:?}", path);
} else {
info!("Loaded peer cache from {:?}", path);
}
Ok(Self {
cache_path: path.to_path_buf(),
})
}
fn read_cache(&self) -> Result<PeerCacheFile> {
if !self.cache_path.exists() {
return Ok(PeerCacheFile::new());
}
let contents = std::fs::read_to_string(&self.cache_path)
.with_context(|| format!("Failed to read cache: {:?}", self.cache_path))?;
serde_json::from_str(&contents)
.context("Failed to parse peer cache JSON")
.or_else(|e| {
warn!("Corrupted peer cache, creating new: {}", e);
Ok(PeerCacheFile::new())
})
}
fn write_cache(&self, cache: &PeerCacheFile) -> Result<()> {
let json = serde_json::to_string_pretty(cache).context("Failed to serialize peer cache")?;
let temp_path = self.cache_path.with_extension("tmp");
std::fs::write(&temp_path, json)
.with_context(|| format!("Failed to write temp cache: {:?}", temp_path))?;
std::fs::rename(&temp_path, &self.cache_path)
.with_context(|| format!("Failed to rename cache: {:?}", self.cache_path))?;
Ok(())
}
pub async fn update_success(&mut self, peer_id: PeerId, addr: SocketAddr) -> Result<()> {
let mut cache = self.read_cache()?;
let peer_id_str = hex::encode(peer_id.as_bytes());
let now = SystemTime::now();
if let Some(entry) = cache.peers.get_mut(&peer_id_str) {
if !entry.addr_hints.contains(&addr) {
entry.addr_hints.push(addr);
}
entry.last_success = now;
entry.success_count += 1;
debug!("Updated peer {} on success", peer_id_str);
} else {
cache.peers.insert(
peer_id_str.clone(),
PeerCacheEntry {
peer_id,
addr_hints: vec![addr],
nat_class: NatClass::Unknown,
roles: Vec::new(),
last_success: now,
success_count: 1,
failure_count: 0,
is_bootstrap: false,
},
);
debug!("Added new peer {} to cache", peer_id_str);
}
self.write_cache(&cache)?;
self.enforce_max_size().await?;
Ok(())
}
pub async fn update_failure(&mut self, peer_id: PeerId) -> Result<()> {
let mut cache = self.read_cache()?;
let peer_id_str = hex::encode(peer_id.as_bytes());
if let Some(entry) = cache.peers.get_mut(&peer_id_str) {
entry.failure_count += 1;
self.write_cache(&cache)?;
debug!("Incremented failure count for peer {}", peer_id_str);
}
Ok(())
}
pub fn get_top_peers(&self, limit: usize) -> Vec<PeerCacheEntry> {
let cache = match self.read_cache() {
Ok(c) => c,
Err(e) => {
warn!("Failed to read peer cache: {}", e);
return Vec::new();
}
};
let mut entries: Vec<PeerCacheEntry> = cache.peers.into_values().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 mut cache = self.read_cache()?;
let initial_count = cache.peers.len();
cache.peers.retain(|_, entry| {
let total_attempts = entry.success_count + entry.failure_count;
if total_attempts <= 5 {
return true; }
let failure_ratio = entry.failure_count as f64 / total_attempts as f64;
failure_ratio <= threshold_ratio
});
let removed = initial_count - cache.peers.len();
if removed > 0 {
self.write_cache(&cache)?;
warn!("Pruned {} failed peers from cache", removed);
}
Ok(())
}
pub fn len(&self) -> usize {
self.read_cache().map(|c| c.peers.len()).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 mut cache = self.read_cache()?;
let peer_id_bytes = blake3::hash(four_words.as_bytes());
let peer_id = PeerId::new(*peer_id_bytes.as_bytes());
let peer_id_str = hex::encode(peer_id.as_bytes());
if !cache.peers.contains_key(&peer_id_str) {
let now = SystemTime::now();
cache.peers.insert(
peer_id_str.clone(),
PeerCacheEntry {
peer_id,
addr_hints: Vec::new(), nat_class: NatClass::Public, roles: vec!["bootstrap".to_string()],
last_success: now,
success_count: 100, failure_count: 0,
is_bootstrap: true,
},
);
self.write_cache(&cache)?;
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 mut cache = self.read_cache()?;
if cache.peers.len() <= Self::MAX_CACHE_SIZE {
return Ok(());
}
let mut non_bootstrap: Vec<(String, SystemTime)> = cache
.peers
.iter()
.filter(|(_, entry)| !entry.is_bootstrap)
.map(|(id, entry)| (id.clone(), entry.last_success))
.collect();
non_bootstrap.sort_by_key(|(_, last_success)| *last_success);
let to_remove = cache.peers.len() - Self::MAX_CACHE_SIZE;
let mut removed = 0;
for (peer_id, _) in non_bootstrap.iter().take(to_remove) {
cache.peers.remove(peer_id);
removed += 1;
}
if removed > 0 {
self.write_cache(&cache)?;
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.json");
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() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("parse addr");
cache
.update_success(peer_id, addr)
.await
.expect("update success");
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
let top_peers = cache.get_top_peers(10);
assert_eq!(top_peers.len(), 1);
assert_eq!(top_peers[0].success_count, 1);
assert_eq!(top_peers[0].failure_count, 0);
}
#[tokio::test]
async fn test_update_failure() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("parse addr");
cache
.update_success(peer_id, addr)
.await
.expect("update success");
cache.update_failure(peer_id).await.expect("update failure");
let top_peers = cache.get_top_peers(10);
assert_eq!(top_peers.len(), 1);
assert_eq!(top_peers[0].failure_count, 1);
}
#[tokio::test]
async fn test_top_peers_sorting() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
for i in 1..=5 {
let peer_id = create_test_peer_id(i);
let addr: SocketAddr = format!("127.0.0.1:808{}", i).parse().expect("parse addr");
for _ in 0..i {
cache
.update_success(peer_id, addr)
.await
.expect("update success");
}
}
let top_peers = cache.get_top_peers(3);
assert_eq!(top_peers.len(), 3);
assert!(top_peers[0].score() >= top_peers[1].score());
assert!(top_peers[1].score() >= top_peers[2].score());
let all_peers = cache.get_top_peers(10);
assert_eq!(all_peers.len(), 5);
}
#[tokio::test]
async fn test_bootstrap_nodes() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
cache
.add_bootstrap_node("ocean-forest-moon-star")
.await
.expect("add bootstrap");
assert_eq!(cache.len(), 1);
let top_peers = cache.get_top_peers(10);
assert_eq!(top_peers.len(), 1);
assert!(top_peers[0].is_bootstrap);
assert_eq!(top_peers[0].success_count, 100); }
#[tokio::test]
async fn test_prune_failed() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
let peer_id = create_test_peer_id(1);
let addr: SocketAddr = "127.0.0.1:8080".parse().expect("parse addr");
cache
.update_success(peer_id, addr)
.await
.expect("update success");
for _ in 0..10 {
cache.update_failure(peer_id).await.expect("update failure");
}
assert_eq!(cache.len(), 1);
cache.prune_failed(0.5).await.expect("prune failed");
assert_eq!(cache.len(), 0);
}
#[tokio::test]
async fn test_max_cache_size() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache = PeerCache::load(&cache_path).await.expect("create cache");
for i in 0..1100 {
let peer_id = create_test_peer_id((i % 256) as u8);
let addr: SocketAddr = format!("127.0.0.1:{}", 8000 + i)
.parse()
.expect("parse addr");
cache
.update_success(peer_id, addr)
.await
.expect("update success");
}
assert!(cache.len() <= PeerCache::MAX_CACHE_SIZE);
}
#[tokio::test]
async fn test_concurrent_writes() {
let temp_dir = TempDir::new().expect("temp dir");
let cache_path = temp_dir.path().join("peers.json");
let mut cache1 = PeerCache::load(&cache_path).await.expect("create cache1");
let mut cache2 = PeerCache::load(&cache_path).await.expect("create cache2");
let peer1 = create_test_peer_id(1);
let peer2 = create_test_peer_id(2);
let addr1: SocketAddr = "127.0.0.1:8081".parse().expect("parse addr1");
let addr2: SocketAddr = "127.0.0.1:8082".parse().expect("parse addr2");
cache1.update_success(peer1, addr1).await.expect("write 1");
cache2.update_success(peer2, addr2).await.expect("write 2");
let cache3 = PeerCache::load(&cache_path).await.expect("create cache3");
assert!(!cache3.is_empty());
}
}