use std::{
collections::HashMap,
fmt,
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use multiaddr::Multiaddr;
#[cfg(feature = "metrics")]
use crate::peer_manager::metrics;
use crate::{
net_address::{MultiaddressesWithStats, PeerAddressSource},
peer_manager::{
NodeId,
PeerFeatures,
PeerFlags,
PeerManagerError,
ThisPeerIdentity,
blocking_storage::BlockingPeerStorage,
peer::Peer,
peer_id::PeerId,
peer_storage_sql::PeerStorageSql,
},
types::{CommsDatabase, CommsPublicKey, TransportProtocol},
};
#[derive(Clone)]
pub struct PeerManager {
peer_storage: BlockingPeerStorage,
transport_protocols: Vec<TransportProtocol>,
ban_generations: Arc<Mutex<HashMap<NodeId, BanGenerationEntry>>>,
}
struct BanGenerationEntry {
generation: u64,
touched_at: Instant,
}
const BAN_GENERATION_TTL: Duration = Duration::from_secs(60 * 60);
const BAN_GENERATION_PRUNE_THRESHOLD: usize = 1024;
const BAN_GENERATION_PRUNE_TARGET: usize = 512;
pub const PEER_LOOKUP_TIMEOUT: Duration = Duration::from_secs(5);
pub const PEER_DATABASE_BUSY_TIMEOUT: Duration = Duration::from_secs(10);
impl PeerManager {
pub fn new(
database: CommsDatabase,
transport_protocols: Vec<TransportProtocol>,
) -> Result<PeerManager, PeerManagerError> {
let peer_storage_sql = PeerStorageSql::new_indexed(database)?;
Ok(Self {
peer_storage: BlockingPeerStorage::new(peer_storage_sql),
transport_protocols,
ban_generations: Arc::new(Mutex::new(HashMap::new())),
})
}
pub fn this_peer_identity(&self) -> ThisPeerIdentity {
self.peer_storage.this_peer_identity()
}
pub async fn count(&self) -> usize {
self.peer_storage.call(|s| Ok(s.count())).await.unwrap_or_default()
}
pub async fn add_or_update_peer(&self, peer: Peer) -> Result<PeerId, PeerManagerError> {
let peer_id = self.peer_storage.call(move |s| s.add_or_update_peer(peer)).await?;
#[cfg(feature = "metrics")]
{
let count = self.count().await;
#[allow(clippy::cast_possible_wrap)]
metrics::peer_list_size().set(count as i64);
}
Ok(peer_id)
}
pub async fn soft_delete_peer(&self, node_id: &NodeId) -> Result<(), PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage.call(move |s| s.soft_delete_peer(&node_id)).await?;
#[cfg(feature = "metrics")]
{
let count = self.count().await;
#[allow(clippy::cast_possible_wrap)]
metrics::peer_list_size().set(count as i64);
}
Ok(())
}
pub async fn get_peers_by_node_ids(&self, node_ids: &[NodeId]) -> Result<Vec<Peer>, PeerManagerError> {
let node_ids = node_ids.to_vec();
self.peer_storage
.call(move |s| s.get_peers_by_node_ids(&node_ids))
.await
}
pub async fn get_peer_public_keys_by_node_ids(
&self,
node_ids: &[NodeId],
) -> Result<Vec<CommsPublicKey>, PeerManagerError> {
let node_ids = node_ids.to_vec();
self.peer_storage
.call(move |s| s.get_peer_public_keys_by_node_ids(&node_ids))
.await
}
pub async fn get_banned_peers(&self) -> Result<Vec<Peer>, PeerManagerError> {
self.peer_storage.call(PeerStorageSql::get_banned_peers).await
}
pub async fn find_by_node_id(&self, node_id: &NodeId) -> Result<Option<Peer>, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage.call(move |s| s.get_peer_by_node_id(&node_id)).await
}
pub async fn get_seed_peers(&self) -> Result<Vec<Peer>, PeerManagerError> {
self.peer_storage.call(PeerStorageSql::get_seed_peers).await
}
pub async fn find_by_public_key(&self, public_key: &CommsPublicKey) -> Result<Option<Peer>, PeerManagerError> {
let public_key = public_key.clone();
self.peer_storage.call(move |s| s.find_by_public_key(&public_key)).await
}
pub async fn find_all_starts_with(&self, partial: &[u8]) -> Result<Vec<Peer>, PeerManagerError> {
let partial = partial.to_vec();
self.peer_storage.call(move |s| s.find_all_starts_with(&partial)).await
}
pub async fn exists(&self, public_key: &CommsPublicKey) -> Result<bool, PeerManagerError> {
let public_key = public_key.clone();
self.peer_storage.call(move |s| s.exists_public_key(&public_key)).await
}
pub async fn exists_node_id(&self, node_id: &NodeId) -> Result<bool, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage.call(move |s| s.exists_node_id(&node_id)).await
}
pub async fn all(&self, features: Option<PeerFeatures>) -> Result<Vec<Peer>, PeerManagerError> {
self.peer_storage.call(move |s| s.all(features)).await
}
pub async fn get_available_dial_candidates(
&self,
exclude_node_ids: &[NodeId],
limit: Option<usize>,
exclude_failed: bool,
randomize: bool,
) -> Result<Vec<Peer>, PeerManagerError> {
let exclude_node_ids = exclude_node_ids.to_vec();
let transport_protocols = self.transport_protocols.clone();
self.peer_storage
.call(move |s| {
s.get_available_dial_candidates(
&exclude_node_ids,
limit,
&transport_protocols,
exclude_failed,
randomize,
)
})
.await
}
pub async fn discovery_syncing(
&self,
n: usize,
excluded_peers: &[NodeId],
features: Option<PeerFeatures>,
external_addresses_only: bool,
max_n: usize,
) -> Result<Vec<Peer>, PeerManagerError> {
let excluded_peers = excluded_peers.to_vec();
self.peer_storage
.call(move |s| s.discovery_syncing(n, &excluded_peers, features, external_addresses_only, max_n))
.await
}
pub async fn add_or_update_online_peer(
&self,
pubkey: &CommsPublicKey,
node_id: &NodeId,
addresses: &[Multiaddr],
peer_features: &PeerFeatures,
source: &PeerAddressSource,
) -> Result<Peer, PeerManagerError> {
let pubkey = pubkey.clone();
let node_id = node_id.clone();
let addresses = addresses.to_vec();
let peer_features = *peer_features;
let source = source.clone();
self.peer_storage
.call(move |s| s.add_or_update_online_peer(&pubkey, &node_id, &addresses, &peer_features, &source))
.await
}
pub async fn direct_identity_node_id(&self, node_id: &NodeId) -> Result<Option<Peer>, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage
.call(move |s| match s.direct_identity_node_id(&node_id) {
Ok(peer) => Ok(Some(peer)),
Err(PeerManagerError::PeerNotFound(_)) | Err(PeerManagerError::BannedPeer) => Ok(None),
Err(err) => Err(err),
})
.await
}
pub async fn direct_identity_public_key(
&self,
public_key: &CommsPublicKey,
) -> Result<Option<Peer>, PeerManagerError> {
let public_key = public_key.clone();
self.peer_storage
.call(move |s| match s.direct_identity_public_key(&public_key) {
Ok(peer) => Ok(Some(peer)),
Err(PeerManagerError::PeerNotFound(_)) | Err(PeerManagerError::BannedPeer) => Ok(None),
Err(err) => Err(err),
})
.await
}
pub async fn get_not_banned_or_deleted_peers(&self) -> Result<Vec<Peer>, PeerManagerError> {
self.peer_storage
.call(PeerStorageSql::get_not_banned_or_deleted_peers)
.await
}
pub async fn random_peers(
&self,
n: usize,
excluded: &[NodeId],
flags: Option<PeerFlags>,
known_good: bool,
) -> Result<Vec<Peer>, PeerManagerError> {
let excluded = excluded.to_vec();
let transport_protocols = self.transport_protocols.clone();
self.peer_storage
.call(move |s| {
let mut peers = s.random_peers(n, &excluded, flags, &transport_protocols, known_good)?;
if known_good && peers.len() < n {
let mut excluded = excluded.clone();
excluded.extend(peers.iter().map(|peer| peer.node_id.clone()));
let mut additional = s.random_peers(
n.checked_sub(peers.len()).unwrap_or(1),
&excluded,
flags,
&transport_protocols,
false,
)?;
peers.append(&mut additional);
}
Ok(peers)
})
.await
}
pub async fn unban_peer(&self, node_id: &NodeId) -> Result<(), PeerManagerError> {
self.bump_ban_generation(node_id);
let node_id = node_id.clone();
self.peer_storage.call(move |s| s.unban_peer(&node_id)).await
}
pub async fn unban_all_peers(&self) -> Result<usize, PeerManagerError> {
self.bump_all_ban_generations();
self.peer_storage.call(PeerStorageSql::unban_all_peers).await
}
pub(crate) fn ban_generation(&self, node_id: &NodeId) -> u64 {
let mut generations = self.ban_generations.lock().unwrap_or_else(|e| e.into_inner());
Self::maybe_prune_ban_generations(&mut generations);
let now = Instant::now();
let entry = generations
.entry(node_id.clone())
.or_insert_with(|| BanGenerationEntry {
generation: 0,
touched_at: now,
});
entry.touched_at = now;
entry.generation
}
pub(crate) fn ban_generation_if_tracked(&self, node_id: &NodeId) -> Option<u64> {
self.ban_generations
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(node_id)
.map(|entry| entry.generation)
}
#[cfg(test)]
pub(crate) fn forget_ban_generation_for_test(&self, node_id: &NodeId) {
self.ban_generations
.lock()
.unwrap_or_else(|e| e.into_inner())
.remove(node_id);
}
fn bump_ban_generation(&self, node_id: &NodeId) {
let mut generations = self.ban_generations.lock().unwrap_or_else(|e| e.into_inner());
Self::maybe_prune_ban_generations(&mut generations);
let now = Instant::now();
let entry = generations
.entry(node_id.clone())
.or_insert_with(|| BanGenerationEntry {
generation: 0,
touched_at: now,
});
entry.generation = entry.generation.saturating_add(1);
entry.touched_at = now;
}
fn bump_all_ban_generations(&self) {
let mut generations = self.ban_generations.lock().unwrap_or_else(|e| e.into_inner());
let now = Instant::now();
for entry in generations.values_mut() {
entry.generation = entry.generation.saturating_add(1);
entry.touched_at = now;
}
}
fn maybe_prune_ban_generations(generations: &mut HashMap<NodeId, BanGenerationEntry>) {
if generations.len() <= BAN_GENERATION_PRUNE_THRESHOLD {
return;
}
let now = Instant::now();
generations.retain(|_, entry| now.saturating_duration_since(entry.touched_at) < BAN_GENERATION_TTL);
if generations.len() <= BAN_GENERATION_PRUNE_TARGET {
return;
}
let excess = generations.len().saturating_sub(BAN_GENERATION_PRUNE_TARGET);
let mut by_age: Vec<(NodeId, Instant)> = generations
.iter()
.map(|(node_id, entry)| (node_id.clone(), entry.touched_at))
.collect();
by_age.sort_by_key(|(_, touched_at)| *touched_at);
for (node_id, _) in by_age.into_iter().take(excess) {
generations.remove(&node_id);
}
}
pub async fn reset_offline_non_wallet_peers(&self) -> Result<usize, PeerManagerError> {
self.peer_storage
.call(PeerStorageSql::reset_offline_non_wallet_peers)
.await
}
pub async fn ban_peer(
&self,
public_key: &CommsPublicKey,
duration: Duration,
reason: String,
) -> Result<NodeId, PeerManagerError> {
let public_key = public_key.clone();
self.peer_storage
.call(move |s| s.ban_peer(&public_key, duration, reason))
.await
}
pub async fn ban_peer_by_node_id(
&self,
node_id: &NodeId,
duration: Duration,
reason: String,
) -> Result<NodeId, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage
.call(move |s| s.ban_peer_by_node_id(&node_id, duration, reason))
.await
}
pub async fn is_peer_banned(&self, node_id: &NodeId) -> Result<bool, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage.call(move |s| s.is_peer_banned(&node_id)).await
}
pub async fn get_peer_features(&self, node_id: &NodeId) -> Result<PeerFeatures, PeerManagerError> {
let peer = self
.find_by_node_id(node_id)
.await?
.ok_or(PeerManagerError::peer_not_found(node_id))?;
Ok(peer.features)
}
pub async fn get_peer_multi_addresses(
&self,
node_id: &NodeId,
) -> Result<MultiaddressesWithStats, PeerManagerError> {
let peer = self
.find_by_node_id(node_id)
.await?
.ok_or(PeerManagerError::peer_not_found(node_id))?;
Ok(peer.addresses)
}
pub async fn get_peers_multi_addresses(
&self,
node_ids: &[NodeId],
) -> Result<Vec<(NodeId, MultiaddressesWithStats)>, PeerManagerError> {
if node_ids.is_empty() {
return Err(PeerManagerError::ProcessError(
"NodeId list cannot be empty".to_string(),
));
}
let peers = self.get_peers_by_node_ids(node_ids).await?;
if peers.is_empty() {
return Err(PeerManagerError::peers_not_found(node_ids));
}
let results = peers.into_iter().map(|p| (p.node_id, p.addresses)).collect::<Vec<_>>();
Ok(results)
}
pub async fn set_peer_metadata(
&self,
node_id: &NodeId,
key: u8,
data: Vec<u8>,
) -> Result<Option<Vec<u8>>, PeerManagerError> {
let node_id = node_id.clone();
self.peer_storage
.call(move |s| s.set_peer_metadata(&node_id, key, data))
.await
}
}
impl fmt::Debug for PeerManager {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("PeerManager { peer_storage: ... }")
}
}
#[cfg(test)]
pub fn create_test_peer(ban_flag: bool, features: PeerFeatures) -> Peer {
use std::borrow::BorrowMut;
use rand::RngExt;
use crate::peer_manager::PeerFlags;
let (_sk, pk) = CommsPublicKey::random_keypair(&mut rand::rng());
let node_id = NodeId::from_key(&pk);
let mut addresses = Vec::new();
for _i in 1..=rand::rng().random_range(1..4) {
let n = [
rand::rng().random_range(11..100),
rand::rng().random_range(1..255),
rand::rng().random_range(1..255),
rand::rng().random_range(1..255),
rand::rng().random_range(5000..9000),
];
let address = format!("/ip4/{}.{}.{}.{}/tcp/{}", n[0], n[1], n[2], n[3], n[4])
.parse::<Multiaddr>()
.unwrap();
addresses.push(address);
}
let net_addresses = MultiaddressesWithStats::from_addresses_with_source(
addresses.clone(),
&create_peer_address_source_with_claim(addresses, features),
);
let mut peer = Peer::new(
pk,
node_id,
net_addresses,
PeerFlags::default(),
features,
Default::default(),
Default::default(),
);
if ban_flag {
peer.ban_for(Duration::from_secs(1000), "".to_string());
}
let good_addresses = peer.addresses.borrow_mut();
let good_address = good_addresses.addresses().first().unwrap().address().clone();
good_addresses.mark_last_seen_now(&good_address);
peer
}
#[cfg(test)]
fn random_onion3_host() -> String {
use rand::distr::Uniform;
const LEN: usize = 56;
const B32: &[u8; 32] = b"abcdefghijklmnopqrstuvwxyz234567";
let mut rng = rand::rng();
let dist = Uniform::new(0, B32.len()).unwrap();
let mut s = String::with_capacity(LEN);
for _ in 0..LEN {
use rand::RngExt;
let idx = rng.sample(dist);
s.push(*B32.get(idx).expect("Index out of bounds") as char);
}
s
}
#[cfg(test)]
pub fn create_test_peer_with_onion_address(ban_flag: bool, features: PeerFeatures) -> Peer {
use std::borrow::BorrowMut;
use rand::RngExt;
use crate::peer_manager::PeerFlags;
let (_sk, pk) = CommsPublicKey::random_keypair(&mut rand::rng());
let node_id = NodeId::from_key(&pk);
let mut addresses = Vec::new();
for _i in 1..=rand::rng().random_range(1..4) {
use std::str::FromStr;
let host = random_onion3_host();
let port = rand::rng().random_range(1024..=65535);
let addr_str = format!("/onion3/{}:{}", host, port);
let address = Multiaddr::from_str(&addr_str).expect("valid onion3 multiaddr");
addresses.push(address);
}
let net_addresses = MultiaddressesWithStats::from_addresses_with_source(
addresses.clone(),
&create_peer_address_source_with_claim(addresses, features),
);
let mut peer = Peer::new(
pk,
node_id,
net_addresses,
PeerFlags::default(),
features,
Default::default(),
Default::default(),
);
if ban_flag {
peer.ban_for(Duration::from_secs(1000), "".to_string());
}
let good_addresses = peer.addresses.borrow_mut();
let good_address = good_addresses.addresses().first().unwrap().address().clone();
good_addresses.mark_last_seen_now(&good_address);
peer
}
#[cfg(test)]
pub fn create_test_peer_add_internal_addresses(ban_flag: bool, features: PeerFeatures) -> Peer {
let mut peer = create_test_peer(ban_flag, features);
add_internal_addresses(&mut peer);
peer
}
#[cfg(test)]
pub fn create_test_peer_internal_addresses_only(ban_flag: bool, features: PeerFeatures) -> Peer {
use crate::peer_manager::PeerFlags;
let (_sk, pk) = CommsPublicKey::random_keypair(&mut rand::rng());
let node_id = NodeId::from_key(&pk);
let mut peer = Peer::new(
pk,
node_id,
MultiaddressesWithStats::default(),
PeerFlags::default(),
features,
Default::default(),
Default::default(),
);
if ban_flag {
peer.ban_for(Duration::from_secs(1000), "".to_string());
}
add_internal_addresses(&mut peer);
peer
}
#[cfg(test)]
fn add_internal_addresses(peer: &mut Peer) {
use rand::{RngExt, prelude::SliceRandom};
let mut addresses = Vec::new();
let address_1 = format!(
"/ip4/127.{}.{}.{}/tcp/{}",
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(9000..9100)
)
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_1);
let address_2 = format!("/ip4/0.0.0.0/tcp/{}", rand::rng().random_range(9100..9200))
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_2);
let address_3 = format!(
"/ip4/10.{}.{}.{}/tcp/{}",
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(9200..9300)
)
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_3);
let address_4 = format!(
"/ip4/172.{}.{}.{}/tcp/{}",
rand::rng().random_range(16..=31),
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(9300..9400)
)
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_4);
let address_5 = format!(
"/ip4/192.168.{}.{}/tcp/{}",
rand::rng().random_range(0..255),
rand::rng().random_range(0..255),
rand::rng().random_range(9400..9500)
)
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_5);
let address_6 = format!("/ip6/::1/tcp/{}", rand::rng().random_range(9500..9600))
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_6);
let address_7 = format!("/ip6/::/tcp/{}", rand::rng().random_range(9600..9700))
.parse::<Multiaddr>()
.unwrap();
addresses.push(address_7);
addresses.shuffle(&mut rand::rng());
peer.addresses
.add_or_update_addresses(&addresses, &PeerAddressSource::Config);
}
#[cfg(test)]
pub fn create_peer_address_source_with_claim(
addresses: Vec<Multiaddr>,
peer_features: PeerFeatures,
) -> PeerAddressSource {
use chrono::Utc;
use tari_crypto::keys::SecretKey;
use crate::{
peer_manager::{IdentitySignature, PeerIdentityClaim},
types::CommsSecretKey,
};
fn create_identity_signature(addresses: &[Multiaddr], peer_features: PeerFeatures) -> IdentitySignature {
let secret = CommsSecretKey::random(&mut rand::rng());
let public_key = CommsPublicKey::from_secret_key(&secret);
let updated_at = Utc::now();
let identity = IdentitySignature::sign_new(&secret, peer_features, addresses, updated_at);
assert!(
identity.is_valid(&public_key, peer_features, addresses).unwrap(),
"Signature is not valid"
);
identity
}
PeerAddressSource::FromPeerConnection {
peer_identity_claim: PeerIdentityClaim {
addresses: addresses.clone(),
features: peer_features,
signature: create_identity_signature(&addresses, peer_features),
},
}
}
#[cfg(test)]
mod test {
#![allow(clippy::indexing_slicing)]
use chrono::{DateTime, Utc};
use tari_common_sqlite::connection::{DbConnection, DbConnectionUrl};
use super::*;
use crate::{
peer_manager::database::{MIGRATIONS, PeerDatabaseSql},
test_utils::node_id,
};
fn create_peer_manager() -> PeerManager {
let db_connection = DbConnection::connect_temp_file_and_migrate(MIGRATIONS).unwrap();
let peers_db = PeerDatabaseSql::new(
db_connection,
&create_test_peer(false, PeerFeatures::COMMUNICATION_NODE),
)
.unwrap();
PeerManager::new(peers_db, TransportProtocol::get_all()).unwrap()
}
fn create_contended_peer_manager() -> (PeerManager, DbConnection, tempfile::TempDir) {
let temp_dir = tempfile::tempdir().unwrap();
let db_url = DbConnectionUrl::File(temp_dir.path().join("contended_peers.db"));
let db_connection =
DbConnection::connect_and_migrate_with_busy_timeout(&db_url, MIGRATIONS, Some(6), Duration::from_secs(5))
.unwrap();
let peers_db = PeerDatabaseSql::new(
db_connection.clone(),
&create_test_peer(false, PeerFeatures::COMMUNICATION_NODE),
)
.unwrap();
let peer_manager = PeerManager::new(peers_db, TransportProtocol::get_all()).unwrap();
(peer_manager, db_connection, temp_dir)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 1)]
async fn peer_database_writes_do_not_park_the_runtime_worker() {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use diesel::connection::SimpleConnection;
let (peer_manager, db_connection, _temp_dir) = create_contended_peer_manager();
let mut lock_holder = db_connection.get_pooled_connection().unwrap();
lock_holder.batch_execute("BEGIN IMMEDIATE;").unwrap();
let progress = Arc::new(AtomicUsize::new(0));
let ticker = tokio::spawn({
let progress = progress.clone();
async move {
loop {
tokio::time::sleep(Duration::from_millis(5)).await;
progress.fetch_add(1, Ordering::Relaxed);
}
}
});
let writers = (0..4)
.map(|_| {
let peer_manager = peer_manager.clone();
let peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
tokio::spawn(async move { peer_manager.add_or_update_peer(peer).await })
})
.collect::<Vec<_>>();
tokio::time::sleep(Duration::from_millis(300)).await;
let ticks = progress.load(Ordering::Relaxed);
assert!(
ticks > 10,
"the runtime worker was parked by the peer database: the canary task only ran {ticks} time(s) in 300ms while 4 peer writes were contended"
);
lock_holder.batch_execute("COMMIT;").unwrap();
drop(lock_holder);
for writer in writers {
writer.await.unwrap().unwrap();
}
ticker.abort();
}
mod storage_access {
use syn::{
Attribute,
Expr,
File,
ImplItem,
ImplItemFn,
Item,
Member,
visit::{self, Visit},
};
fn is_cfg_test(attrs: &[Attribute]) -> bool {
attrs.iter().any(|attr| {
if !attr.path().is_ident("cfg") {
return false;
}
let mut mentions_test = false;
let _result = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("test") {
mentions_test = true;
}
Ok(())
});
mentions_test
})
}
fn item_attrs(item: &Item) -> &[Attribute] {
match item {
Item::Const(i) => &i.attrs,
Item::Enum(i) => &i.attrs,
Item::Fn(i) => &i.attrs,
Item::Impl(i) => &i.attrs,
Item::Macro(i) => &i.attrs,
Item::Mod(i) => &i.attrs,
Item::Static(i) => &i.attrs,
Item::Struct(i) => &i.attrs,
Item::Trait(i) => &i.attrs,
Item::Type(i) => &i.attrs,
Item::Union(i) => &i.attrs,
Item::Use(i) => &i.attrs,
_ => &[],
}
}
#[derive(Debug, Default)]
pub(super) struct FieldUsage {
pub(super) accesses: usize,
pub(super) methods: Vec<String>,
}
struct FieldVisitor<'a> {
field: &'a str,
usage: FieldUsage,
}
fn is_self_field(expr: &Expr, field: &str) -> bool {
let Expr::Field(access) = expr else {
return false;
};
let Expr::Path(base) = access.base.as_ref() else {
return false;
};
let Member::Named(name) = &access.member else {
return false;
};
base.path.is_ident("self") && name == field
}
impl<'ast> Visit<'ast> for FieldVisitor<'_> {
fn visit_item(&mut self, item: &'ast Item) {
if is_cfg_test(item_attrs(item)) {
return;
}
visit::visit_item(self, item);
}
fn visit_impl_item(&mut self, item: &'ast ImplItem) {
let attrs = match item {
ImplItem::Const(i) => &i.attrs,
ImplItem::Fn(i) => &i.attrs,
ImplItem::Type(i) => &i.attrs,
ImplItem::Macro(i) => &i.attrs,
_ => &[][..],
};
if is_cfg_test(attrs) {
return;
}
visit::visit_impl_item(self, item);
}
fn visit_expr(&mut self, expr: &'ast Expr) {
match expr {
Expr::MethodCall(call) if is_self_field(&call.receiver, self.field) => {
self.usage.methods.push(call.method.to_string());
},
_ => {},
}
if is_self_field(expr, self.field) {
self.usage.accesses = self.usage.accesses.saturating_add(1);
}
visit::visit_expr(self, expr);
}
}
pub(super) fn field_usage(file: &File, field: &str) -> FieldUsage {
let mut visitor = FieldVisitor {
field,
usage: FieldUsage::default(),
};
visitor.visit_file(file);
visitor.usage
}
pub(super) fn inherent_method<'a>(file: &'a File, type_name: &str, method: &str) -> Option<&'a ImplItemFn> {
file.items.iter().find_map(|item| {
let Item::Impl(block) = item else {
return None;
};
if block.trait_.is_some() || is_cfg_test(&block.attrs) {
return None;
}
let syn::Type::Path(path) = block.self_ty.as_ref() else {
return None;
};
if path.path.segments.last().is_none_or(|seg| seg.ident != type_name) {
return None;
}
block.items.iter().find_map(|impl_item| match impl_item {
ImplItem::Fn(func) if func.sig.ident == method => Some(func),
_ => None,
})
})
}
struct CallCounter<'a> {
callee: &'a str,
count: usize,
}
impl<'ast> Visit<'ast> for CallCounter<'_> {
fn visit_expr(&mut self, expr: &'ast Expr) {
if let Expr::Call(call) = expr &&
let Expr::Path(path) = call.func.as_ref() &&
path.path.is_ident(self.callee)
{
self.count = self.count.saturating_add(1);
}
visit::visit_expr(self, expr);
}
}
fn count_calls_in_block(block: &syn::Block, callee: &str) -> usize {
let mut counter = CallCounter { callee, count: 0 };
counter.visit_block(block);
counter.count
}
fn count_calls_in_expr(expr: &Expr, callee: &str) -> usize {
let mut counter = CallCounter { callee, count: 0 };
counter.visit_expr(expr);
counter.count
}
pub(super) fn dispatch_balance(func: &ImplItemFn, callee: &str, dispatcher: &str) -> (usize, usize) {
let total = count_calls_in_block(&func.block, callee);
struct DispatcherVisitor<'a> {
dispatcher: &'a str,
callee: &'a str,
dispatched: usize,
}
impl<'ast> Visit<'ast> for DispatcherVisitor<'_> {
fn visit_expr(&mut self, expr: &'ast Expr) {
if let Expr::Call(call) = expr &&
let Expr::Path(path) = call.func.as_ref() &&
path.path
.segments
.last()
.is_some_and(|seg| seg.ident == self.dispatcher) &&
let Some(Expr::Closure(closure)) = call.args.first()
{
self.dispatched = self
.dispatched
.saturating_add(count_calls_in_expr(closure.body.as_ref(), self.callee));
}
visit::visit_expr(self, expr);
}
}
let mut visitor = DispatcherVisitor {
dispatcher,
callee,
dispatched: 0,
};
visitor.visit_block(&func.block);
(total, visitor.dispatched)
}
}
fn assert_field_only_reached_via(source: &str, file: &str, field: &str, permitted: &[&str]) -> usize {
let parsed = syn::parse_file(source).unwrap_or_else(|err| panic!("{file} must parse: {err}"));
let usage = storage_access::field_usage(&parsed, field);
let offending = usage
.methods
.iter()
.filter(|method| !permitted.contains(&method.as_str()))
.cloned()
.collect::<Vec<_>>();
assert!(
offending.is_empty(),
"{file}: `self.{field}` is used as the receiver of {offending:?}, but only {permitted:?} are permitted. \
Anything else runs SQLite on the caller's tokio worker thread."
);
assert_eq!(
usage.accesses,
usage.methods.len(),
"{file}: {} of the {} `self.{field}` access(es) are not direct method calls. The handle must not escape - \
once it does, nothing constrains how it is used.",
usage.accesses.saturating_sub(usage.methods.len()),
usage.accesses,
);
usage.accesses
}
#[test]
fn no_peer_manager_method_performs_synchronous_storage_io() {
const MANAGER_SOURCE: &str = include_str!("manager.rs");
const BLOCKING_STORAGE_SOURCE: &str = include_str!("blocking_storage.rs");
let inspected = assert_field_only_reached_via(MANAGER_SOURCE, "manager.rs", "peer_storage", &[
"call",
"this_peer_identity",
]);
assert!(
inspected >= 20,
"manager.rs: the guard found only {inspected} `self.peer_storage` access(es), so it is no longer guarding \
anything meaningful. Check that the field name still matches the source."
);
let inspected = assert_field_only_reached_via(BLOCKING_STORAGE_SOURCE, "blocking_storage.rs", "storage", &[
"clone",
"this_peer_identity",
]);
assert!(
inspected > 0,
"blocking_storage.rs: the guard found no `self.storage` accesses at all, so the field name no longer \
matches the source and this guard is inert."
);
let parsed = syn::parse_file(BLOCKING_STORAGE_SOURCE).expect("blocking_storage.rs must parse");
let call = storage_access::inherent_method(&parsed, "BlockingPeerStorage", "call")
.expect("blocking_storage.rs must declare `BlockingPeerStorage::call`");
let (total, dispatched) = storage_access::dispatch_balance(call, "f", "spawn_blocking");
assert!(
total > 0,
"`BlockingPeerStorage::call` never invokes the caller's closure `f`; this guard is inert."
);
assert_eq!(
total, dispatched,
"`BlockingPeerStorage::call` invokes the caller's closure `f` {total} time(s) but only {dispatched} of \
those are inside a closure passed to `spawn_blocking`. The rest run SQLite on the caller's tokio worker \
thread."
);
}
#[tokio::test]
#[allow(clippy::too_many_lines)]
async fn test_get_broadcast_identities() {
let peer_manager = create_peer_manager();
let mut test_peers = vec![create_test_peer(true, PeerFeatures::COMMUNICATION_NODE)];
assert!(
peer_manager
.add_or_update_peer(test_peers[test_peers.len() - 1].clone())
.await
.is_ok()
);
for _i in 0..18 {
test_peers.push(create_test_peer(false, PeerFeatures::COMMUNICATION_NODE));
assert!(
peer_manager
.add_or_update_peer(test_peers[test_peers.len() - 1].clone())
.await
.is_ok()
);
}
test_peers.push(create_test_peer(true, PeerFeatures::COMMUNICATION_NODE));
assert!(
peer_manager
.add_or_update_peer(test_peers[test_peers.len() - 1].clone())
.await
.is_ok()
);
let selected_peers = peer_manager
.direct_identity_node_id(&test_peers[2].node_id)
.await
.unwrap()
.unwrap();
assert_eq!(selected_peers.node_id, test_peers[2].node_id);
assert_eq!(selected_peers.public_key, test_peers[2].public_key);
let unmanaged_peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
assert!(
peer_manager
.direct_identity_node_id(&unmanaged_peer.node_id)
.await
.unwrap()
.is_none()
);
let selected_peers = peer_manager.get_not_banned_or_deleted_peers().await.unwrap();
assert_eq!(selected_peers.len(), 18);
for peer_identity in &selected_peers {
assert!(
!peer_manager
.find_by_node_id(&peer_identity.node_id)
.await
.unwrap()
.unwrap()
.is_banned(),
);
}
let identities1 = peer_manager.random_peers(10, &[], None, false).await.unwrap();
let identities2 = peer_manager.random_peers(10, &[], None, false).await.unwrap();
assert_ne!(identities1, identities2);
}
#[tokio::test]
async fn test_add_or_update_online_peer() {
let peer_manager = create_peer_manager();
let peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
peer_manager.add_or_update_peer(peer.clone()).await.unwrap();
let peer = peer_manager
.add_or_update_online_peer(
&peer.public_key,
&peer.node_id,
&[],
&peer.features,
&PeerAddressSource::Config,
)
.await
.unwrap();
assert!(!peer.is_offline());
}
async fn validate_claim_bump_by_newer(
peer_manager: &PeerManager,
update_peer: &Peer,
previous_claim_time: Option<DateTime<Utc>>,
expected_count: usize,
) -> DateTime<Utc> {
let peer_from_db = peer_manager
.find_by_node_id(&update_peer.node_id)
.await
.unwrap()
.unwrap();
let newest_time = peer_from_db.addresses.newest_claim_updated_at().unwrap();
if let Some(prev_time) = previous_claim_time {
assert!(
newest_time > prev_time,
"New claim time was not newer than previous claim time"
);
}
for addr in peer_from_db.addresses.addresses() {
let claim_time = match addr.source() {
PeerAddressSource::FromPeerConnection { peer_identity_claim } => {
peer_identity_claim.signature.updated_at()
},
_ => panic!("Expected FromPeerConnection source for address: {}", addr.address()),
};
assert_eq!(
claim_time, newest_time,
"Address claim time inconsistent among addresses"
);
}
assert_eq!(peer_manager.count().await, expected_count, "Peer count mismatch");
let mut expected_addresses = update_peer
.addresses
.addresses()
.iter()
.map(|a| a.address().clone())
.collect::<Vec<_>>();
let mut addresses_from_db = peer_from_db
.addresses
.addresses()
.iter()
.map(|a| a.address().clone())
.collect::<Vec<_>>();
expected_addresses.sort();
addresses_from_db.sort();
assert_eq!(expected_addresses, addresses_from_db);
newest_time
}
#[tokio::test]
async fn it_correctly_merges_old_and_new_address_claims() {
let peer_manager = create_peer_manager();
let mut peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
peer_manager.add_or_update_peer(peer.clone()).await.unwrap();
let claim_1_time = validate_claim_bump_by_newer(&peer_manager, &peer, None, 1).await;
tokio::time::sleep(Duration::from_millis(150)).await; let peer_addresses = peer
.addresses
.addresses()
.iter()
.map(|a| a.address().clone())
.collect::<Vec<_>>();
let peer_address_source = create_peer_address_source_with_claim(peer_addresses.clone(), peer.features);
peer.addresses
.add_or_update_addresses(&peer_addresses, &peer_address_source);
peer_manager.add_or_update_peer(peer.clone()).await.unwrap();
let claim_2_time = validate_claim_bump_by_newer(&peer_manager, &peer, Some(claim_1_time), 1).await;
tokio::time::sleep(Duration::from_millis(150)).await; let mut update_peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
update_peer.node_id = peer.node_id.clone();
update_peer.public_key = peer.public_key.clone();
peer_manager.add_or_update_peer(update_peer.clone()).await.unwrap();
let _claim_3_time = validate_claim_bump_by_newer(&peer_manager, &update_peer, Some(claim_2_time), 1).await;
peer_manager.add_or_update_peer(peer.clone()).await.unwrap();
let _claim_3_time = validate_claim_bump_by_newer(&peer_manager, &update_peer, Some(claim_2_time), 1).await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
async fn test_concurrent_add_or_update_and_get_random_peers() {
let peer_manager = create_peer_manager();
let num_peers = 75;
let num_write_tasks = 20;
let num_read_tasks = 1500;
let n = 100;
let add_tasks: Vec<_> = (0..num_write_tasks)
.map(|_| {
let peer_manager = peer_manager.clone();
tokio::spawn(async move {
let mut peers_to_update_last_seen = Vec::new();
let mut peers_to_set_metadata = Vec::new();
for i in 0..num_peers {
let peer = create_test_peer(false, PeerFeatures::COMMUNICATION_NODE);
if i % 7 == 0 {
peers_to_update_last_seen.push(peer.clone());
}
if i % 11 == 0 {
peers_to_set_metadata.push(peer.clone());
}
peers_to_update_last_seen.push(peer.clone());
peer_manager.add_or_update_peer(peer).await.unwrap();
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
}
for peer in &mut peers_to_update_last_seen {
let addresses = peer.addresses.addresses().to_vec();
peer.addresses.mark_last_seen_now(addresses[0].address());
peer_manager.add_or_update_peer(peer.clone()).await.unwrap();
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
}
for (key, peer) in peers_to_set_metadata.iter().enumerate() {
peer_manager
.set_peer_metadata(
&peer.node_id,
u8::try_from(key % usize::from(u8::MAX)).unwrap_or_default(),
vec![1, 2, 3],
)
.await
.unwrap();
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
}
Ok::<_, PeerManagerError>(())
})
})
.collect();
let get_tasks: Vec<_> = (0..num_read_tasks)
.map(|_| {
let peer_manager = peer_manager.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
let _random_peers = peer_manager.random_peers(n, &[], None, false).await.unwrap();
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
let _total_peers = peer_manager.count().await;
Ok::<_, PeerManagerError>(())
})
})
.collect();
let all_tasks = add_tasks.into_iter().chain(get_tasks);
for (i, task) in all_tasks.enumerate() {
match task.await {
Ok(Ok(_)) => { },
Ok(Err(e)) => panic!("Task {i} failed with PeerManagerError: {e:?}"),
Err(e) => panic!("Task {i} panicked: {e:?}"),
}
}
tokio::time::sleep(Duration::from_micros(rand::random::<u64>() % 100)).await;
let random_peers = peer_manager.random_peers(n, &[], None, false).await.unwrap();
let total_peers = peer_manager.count().await;
assert_eq!(total_peers, num_peers * num_write_tasks);
assert!(random_peers.len() <= n);
}
#[tokio::test]
async fn unban_peer_bumps_the_ban_generation() {
let peer_manager = create_peer_manager();
let node_id = node_id::random();
let other_node_id = node_id::random();
let baseline = peer_manager.ban_generation(&node_id);
let other_baseline = peer_manager.ban_generation(&other_node_id);
peer_manager.unban_peer(&node_id).await.unwrap();
assert_ne!(
peer_manager.ban_generation(&node_id),
baseline,
"unban_peer must bump the generation so a pending retry snapshotted before it can detect the change"
);
assert_eq!(
peer_manager.ban_generation(&other_node_id),
other_baseline,
"unban_peer must not bump any other peer's generation"
);
}
#[tokio::test]
async fn ban_generation_if_tracked_never_reads_a_missing_entry_as_unchanged() {
let peer_manager = create_peer_manager();
let node_id = node_id::random();
assert_eq!(peer_manager.ban_generation_if_tracked(&node_id), None);
let expected_generation = peer_manager.ban_generation(&node_id);
assert_eq!(
expected_generation, 0,
"a peer's first-ever baseline is 0 - the exact value that made the bug possible"
);
assert_eq!(
peer_manager.ban_generation_if_tracked(&node_id),
Some(expected_generation)
);
peer_manager.ban_generations.lock().unwrap().remove(&node_id);
assert_ne!(
peer_manager.ban_generation_if_tracked(&node_id),
Some(expected_generation),
"a missing entry must never compare equal to a captured baseline, even when that baseline is 0"
);
assert_eq!(
peer_manager.ban_generation_if_tracked(&node_id),
None,
"ban_generation_if_tracked must never re-insert on a miss - calling it must not be observable as a write"
);
}
#[tokio::test]
async fn unban_all_peers_bumps_every_tracked_generation() {
let peer_manager = create_peer_manager();
let read_only_node_id = node_id::random();
let previously_unbanned_node_id = node_id::random();
let read_only_baseline = peer_manager.ban_generation(&read_only_node_id);
peer_manager.unban_peer(&previously_unbanned_node_id).await.unwrap();
let previously_unbanned_baseline = peer_manager.ban_generation(&previously_unbanned_node_id);
peer_manager.unban_all_peers().await.unwrap();
assert_ne!(
peer_manager.ban_generation(&read_only_node_id),
read_only_baseline,
"unban_all_peers must bump a peer's generation even if it was only ever read, never individually unbanned before"
);
assert_ne!(
peer_manager.ban_generation(&previously_unbanned_node_id),
previously_unbanned_baseline,
"unban_all_peers must bump a peer's generation again even if it was already bumped once before"
);
}
#[test]
fn maybe_prune_ban_generations_evicts_oldest_down_to_target_when_over_threshold() {
let now = Instant::now();
let total = BAN_GENERATION_PRUNE_THRESHOLD + 5;
let mut generations: HashMap<NodeId, BanGenerationEntry> = HashMap::with_capacity(total);
let mut by_age = Vec::with_capacity(total);
for i in 0..total {
let node_id = node_id::random();
let touched_at = now.checked_sub(Duration::from_millis(i as u64)).unwrap_or(now);
generations.insert(node_id.clone(), BanGenerationEntry {
generation: 0,
touched_at,
});
by_age.push(node_id);
}
by_age.reverse();
PeerManager::maybe_prune_ban_generations(&mut generations);
assert_eq!(
generations.len(),
BAN_GENERATION_PRUNE_TARGET,
"a sweep that needed the oldest-eviction stage must bring the map down to exactly the target, not merely under the threshold"
);
let evicted_count = total - BAN_GENERATION_PRUNE_TARGET;
for node_id in &by_age[..evicted_count] {
assert!(
!generations.contains_key(node_id),
"one of the oldest entries was not evicted"
);
}
for node_id in &by_age[evicted_count..] {
assert!(
generations.contains_key(node_id),
"a newer entry was evicted while an older one survived"
);
}
}
}