use std::net::SocketAddr;
use std::sync::atomic::{AtomicU32, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use chia_protocol::{Message, NewPeakWallet, ProtocolMessageTypes};
use chia_traits::Streamable;
use futures_util::stream::{FuturesUnordered, StreamExt};
use tokio::sync::{mpsc, RwLock};
use chia_wallet_sdk::client::Peer;
use tokio_tungstenite::Connector;
use crate::types::ChiaQueryError;
use crate::NetworkType;
use super::connect;
struct PeerEntry {
peer: Peer,
address: SocketAddr,
origin: connect::PeerOrigin,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PeerRequirement {
Required,
Optional,
}
pub struct PeerPool {
entries: RwLock<Vec<PeerEntry>>,
next_idx: AtomicUsize,
max_peers: usize,
tls: Connector,
network: NetworkType,
connect_timeout: Duration,
peak_height: Arc<AtomicU32>,
}
impl PeerPool {
pub async fn new(
network: NetworkType,
tls: Connector,
max_peers: usize,
requirement: PeerRequirement,
connect_timeout: Duration,
) -> Result<Self, ChiaQueryError> {
let peak_height = Arc::new(AtomicU32::new(0));
let mut futures = FuturesUnordered::new();
for _ in 0..max_peers {
let t = tls.clone();
futures.push(async move {
connect::connect_random_peer_excluding(network, &t, connect_timeout, &[]).await
});
}
let mut connected = Vec::new();
while let Some(result) = futures.next().await {
match result {
Ok(connection) => connected.push(connection),
Err(e) => log::debug!("initial peer connect failed: {e}"),
}
}
let pool = Self {
entries: RwLock::new(Vec::new()),
next_idx: AtomicUsize::new(0),
max_peers,
tls,
network,
connect_timeout,
peak_height,
};
for (peer, addr, receiver, origin) in connected {
if pool.admit(peer, addr, origin).await {
pool.spawn_receiver_handler(receiver);
}
}
if !pool.has_peers().await {
if requirement == PeerRequirement::Required {
return Err(ChiaQueryError::PeerDiscoveryFailed);
}
log::warn!("no peers connected; serving from the coinset fallback until one does");
}
Ok(pool)
}
pub fn peak_height(&self) -> u32 {
self.peak_height.load(Ordering::Relaxed)
}
pub async fn select_peer(&self) -> Option<(Peer, SocketAddr)> {
let entries = self.entries.read().await;
if entries.is_empty() {
return None;
}
let idx = self.next_idx.fetch_add(1, Ordering::Relaxed) % entries.len();
let entry = &entries[idx];
Some((entry.peer.clone(), entry.address))
}
pub async fn select_corroborating_peer(&self, asked: SocketAddr) -> Option<(Peer, SocketAddr)> {
let entries = self.entries.read().await;
let candidates: Vec<&PeerEntry> = entries
.iter()
.filter(|e| e.address != asked && e.origin == connect::PeerOrigin::Discovered)
.collect();
if candidates.is_empty() {
return None;
}
let idx = self.next_idx.fetch_add(1, Ordering::Relaxed) % candidates.len();
let entry = candidates[idx];
Some((entry.peer.clone(), entry.address))
}
pub async fn eject_peer(&self, addr: SocketAddr) {
{
let mut entries = self.entries.write().await;
entries.retain(|e| e.address != addr);
}
log::debug!(
"peer ejected from pool; will refill on next request (network={:?})",
self.network,
);
}
pub async fn has_peers(&self) -> bool {
!self.entries.read().await.is_empty()
}
pub async fn peer_count(&self) -> usize {
self.entries.read().await.len()
}
pub async fn independent_peer_count(&self) -> usize {
self.entries
.read()
.await
.iter()
.filter(|e| e.origin == connect::PeerOrigin::Discovered)
.count()
}
async fn admit(&self, peer: Peer, address: SocketAddr, origin: connect::PeerOrigin) -> bool {
let mut entries = self.entries.write().await;
if entries.len() >= self.max_peers {
log::debug!("peer {address} not admitted: pool is at capacity");
return false;
}
if entries.iter().any(|e| e.address == address) {
log::debug!("peer {address} not admitted: already held");
return false;
}
entries.push(PeerEntry {
peer,
address,
origin,
});
log::debug!("peer admitted: {address} ({origin:?})");
true
}
pub async fn try_refill(&self) {
let held: Vec<SocketAddr> = {
let entries = self.entries.read().await;
if entries.len() >= self.max_peers {
return;
}
entries.iter().map(|e| e.address).collect()
};
match connect::connect_random_peer_excluding(
self.network,
&self.tls,
self.connect_timeout,
&held,
)
.await
{
Ok((peer, addr, receiver, origin)) => {
if self.admit(peer, addr, origin).await {
self.spawn_receiver_handler(receiver);
log::debug!("replacement peer connected: {addr}");
}
}
Err(e) => log::warn!("replacement peer connect failed: {e}"),
}
}
pub fn spawn_receiver_handler(&self, mut receiver: mpsc::Receiver<Message>) {
let peak = Arc::clone(&self.peak_height);
tokio::spawn(async move {
while let Some(msg) = receiver.recv().await {
if msg.msg_type == ProtocolMessageTypes::NewPeakWallet {
if let Ok(new_peak) = NewPeakWallet::from_bytes(&msg.data) {
let prev = peak.fetch_max(new_peak.height, Ordering::Relaxed);
if new_peak.height > prev {
log::debug!("new peak from peer: {}", new_peak.height);
}
}
}
}
});
}
}
#[cfg(test)]
impl PeerPool {
pub(crate) fn for_tests(max_peers: usize) -> Self {
Self {
entries: RwLock::new(Vec::new()),
next_idx: AtomicUsize::new(0),
max_peers,
tls: connect::create_generated_tls().expect("generate a TLS identity"),
network: NetworkType::Mainnet,
connect_timeout: Duration::from_millis(1),
peak_height: Arc::new(AtomicU32::new(0)),
}
}
pub(crate) async fn admit_for_tests(
&self,
peer: Peer,
address: SocketAddr,
origin: connect::PeerOrigin,
) -> bool {
self.admit(peer, address, origin).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::peer::connect::{create_generated_tls, PeerOrigin};
use crate::peer::test_support::{address, loopback_peer};
use super::PeerPool as _Pool;
fn empty_pool(max_peers: usize) -> PeerPool {
_Pool::for_tests(max_peers)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn one_address_cannot_fill_the_pool_however_many_fills_race() {
let pool = Arc::new(empty_pool(8));
let peer = loopback_peer().await;
let occupied = address(1);
let mut fills = Vec::new();
for _ in 0..8 {
let pool = Arc::clone(&pool);
let peer = peer.clone();
fills.push(tokio::spawn(async move {
pool.admit(peer, occupied, PeerOrigin::Priority).await
}));
}
let admitted = futures_util::future::join_all(fills)
.await
.into_iter()
.filter(|r| *r.as_ref().expect("the admission task must not panic"))
.count();
assert_eq!(
admitted, 1,
"exactly one fill of an address may be admitted"
);
assert_eq!(
pool.peer_count().await,
1,
"eight concurrent fills of one address must leave one connection, not eight"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn distinct_addresses_all_fill_concurrently() {
let pool = Arc::new(empty_pool(8));
let peer = loopback_peer().await;
let mut fills = Vec::new();
for octet in 1..=8u8 {
let pool = Arc::clone(&pool);
let peer = peer.clone();
fills.push(tokio::spawn(async move {
pool.admit(peer, address(octet), PeerOrigin::Discovered)
.await
}));
}
futures_util::future::join_all(fills).await;
assert_eq!(
pool.peer_count().await,
8,
"eight distinct addresses must all be admitted"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_fills_never_exceed_max_peers() {
let pool = Arc::new(empty_pool(3));
let peer = loopback_peer().await;
let mut fills = Vec::new();
for octet in 1..=10u8 {
let pool = Arc::clone(&pool);
let peer = peer.clone();
fills.push(tokio::spawn(async move {
pool.admit(peer, address(octet), PeerOrigin::Discovered)
.await
}));
}
futures_util::future::join_all(fills).await;
assert_eq!(pool.peer_count().await, 3, "max_peers is a hard ceiling");
}
#[tokio::test]
async fn a_preferred_peer_is_held_but_not_counted_as_an_independent_opinion() {
let pool = empty_pool(5);
let peer = loopback_peer().await;
assert!(
pool.admit(peer.clone(), address(1), PeerOrigin::Priority)
.await
);
assert!(
pool.admit(peer.clone(), address(2), PeerOrigin::Discovered)
.await
);
assert!(pool.admit(peer, address(3), PeerOrigin::Discovered).await);
assert_eq!(pool.peer_count().await, 3, "three connections are held");
assert_eq!(
pool.independent_peer_count().await,
2,
"the co-resident peer is held and read from, but is not an independent voice"
);
}
#[tokio::test]
async fn an_ejected_address_can_be_admitted_again() {
let pool = empty_pool(5);
let peer = loopback_peer().await;
let addr = address(1);
assert!(pool.admit(peer.clone(), addr, PeerOrigin::Discovered).await);
assert!(
!pool.admit(peer.clone(), addr, PeerOrigin::Discovered).await,
"still held, so still a duplicate"
);
pool.eject_peer(addr).await;
assert!(
pool.admit(peer, addr, PeerOrigin::Discovered).await,
"a re-dialled peer must be admissible after ejection"
);
assert_eq!(pool.peer_count().await, 1);
}
async fn pool_with_no_connection_attempts(
requirement: PeerRequirement,
) -> Result<PeerPool, ChiaQueryError> {
PeerPool::new(
NetworkType::Mainnet,
create_generated_tls().expect("generate a TLS identity"),
0,
requirement,
Duration::from_millis(1),
)
.await
}
#[tokio::test]
async fn empty_pool_is_fatal_when_peers_are_required() {
assert!(matches!(
pool_with_no_connection_attempts(PeerRequirement::Required).await,
Err(ChiaQueryError::PeerDiscoveryFailed)
));
}
#[tokio::test]
async fn empty_pool_is_tolerated_when_peers_are_optional() {
let pool = pool_with_no_connection_attempts(PeerRequirement::Optional)
.await
.expect("an optional peer pool must construct with zero peers");
assert!(!pool.has_peers().await);
}
#[tokio::test]
async fn an_unfilled_pool_counts_what_it_holds_not_the_target_it_was_given() {
let pool = empty_pool(5);
assert_eq!(
pool.peer_count().await,
0,
"held is 0 while the target is 5"
);
assert!(!pool.has_peers().await);
}
}