1use std::net::SocketAddr;
2use std::sync::atomic::{AtomicU32, AtomicUsize, Ordering};
3use std::sync::Arc;
4use std::time::Duration;
5
6use chia_protocol::{Message, NewPeakWallet, ProtocolMessageTypes};
7use chia_traits::Streamable;
8use futures_util::stream::{FuturesUnordered, StreamExt};
9use tokio::sync::{mpsc, RwLock};
10
11use chia_wallet_sdk::client::Peer;
12use tokio_tungstenite::Connector;
13
14use crate::types::ChiaQueryError;
15use crate::NetworkType;
16
17use super::connect;
18
19struct PeerEntry {
24 peer: Peer,
25 address: SocketAddr,
26}
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub enum PeerRequirement {
39 Required,
41 Optional,
43}
44
45pub struct PeerPool {
50 entries: RwLock<Vec<PeerEntry>>,
51 next_idx: AtomicUsize,
52 max_peers: usize,
53 tls: Connector,
54 network: NetworkType,
55 connect_timeout: Duration,
56 peak_height: Arc<AtomicU32>,
59}
60
61impl PeerPool {
62 pub async fn new(
67 network: NetworkType,
68 tls: Connector,
69 max_peers: usize,
70 requirement: PeerRequirement,
71 connect_timeout: Duration,
72 ) -> Result<Self, ChiaQueryError> {
73 let peak_height = Arc::new(AtomicU32::new(0));
74
75 let mut futures = FuturesUnordered::new();
77 for _ in 0..max_peers {
78 let t = tls.clone();
79 futures.push(async move {
80 connect::connect_random_peer(network, &t, connect_timeout).await
81 });
82 }
83
84 let mut initial: Vec<PeerEntry> = Vec::new();
85 let mut receivers = Vec::new();
86 while let Some(result) = futures.next().await {
87 match result {
88 Ok((peer, addr, receiver)) => {
89 initial.push(PeerEntry {
90 peer,
91 address: addr,
92 });
93 receivers.push(receiver);
94 }
95 Err(e) => log::debug!("initial peer connect failed: {e}"),
96 }
97 }
98
99 if initial.is_empty() {
100 if requirement == PeerRequirement::Required {
101 return Err(ChiaQueryError::PeerDiscoveryFailed);
102 }
103 log::warn!("no peers connected; serving from the coinset fallback until one does");
104 }
105
106 let pool = Self {
107 entries: RwLock::new(initial),
108 next_idx: AtomicUsize::new(0),
109 max_peers,
110 tls,
111 network,
112 connect_timeout,
113 peak_height,
114 };
115
116 for receiver in receivers {
119 pool.spawn_receiver_handler(receiver);
120 }
121
122 Ok(pool)
123 }
124
125 pub fn peak_height(&self) -> u32 {
128 self.peak_height.load(Ordering::Relaxed)
129 }
130
131 pub async fn select_peer(&self) -> Option<(Peer, SocketAddr)> {
134 let entries = self.entries.read().await;
135 if entries.is_empty() {
136 return None;
137 }
138 let idx = self.next_idx.fetch_add(1, Ordering::Relaxed) % entries.len();
139 let entry = &entries[idx];
140 Some((entry.peer.clone(), entry.address))
141 }
142
143 pub async fn eject_peer(&self, addr: SocketAddr) {
145 {
146 let mut entries = self.entries.write().await;
147 entries.retain(|e| e.address != addr);
148 }
149 log::debug!(
150 "peer ejected from pool; will refill on next request (network={:?})",
151 self.network,
152 );
153 }
154
155 pub async fn has_peers(&self) -> bool {
157 !self.entries.read().await.is_empty()
158 }
159
160 pub async fn try_refill(&self) {
164 let current = self.entries.read().await.len();
165 if current >= self.max_peers {
166 return;
167 }
168 match connect::connect_random_peer(self.network, &self.tls, self.connect_timeout).await {
169 Ok((peer, addr, receiver)) => {
170 self.spawn_receiver_handler(receiver);
171 let mut entries = self.entries.write().await;
172 if entries.len() < self.max_peers {
173 entries.push(PeerEntry {
174 peer,
175 address: addr,
176 });
177 log::debug!("replacement peer connected: {addr}");
178 }
179 }
180 Err(e) => log::warn!("replacement peer connect failed: {e}"),
181 }
182 }
183
184 pub fn spawn_receiver_handler(&self, mut receiver: mpsc::Receiver<Message>) {
192 let peak = Arc::clone(&self.peak_height);
193 tokio::spawn(async move {
194 while let Some(msg) = receiver.recv().await {
195 if msg.msg_type == ProtocolMessageTypes::NewPeakWallet {
196 if let Ok(new_peak) = NewPeakWallet::from_bytes(&msg.data) {
197 let prev = peak.fetch_max(new_peak.height, Ordering::Relaxed);
198 if new_peak.height > prev {
199 log::debug!("new peak from peer: {}", new_peak.height);
200 }
201 }
202 }
203 }
204 });
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use crate::peer::connect::create_generated_tls;
212
213 async fn pool_with_no_connection_attempts(
216 requirement: PeerRequirement,
217 ) -> Result<PeerPool, ChiaQueryError> {
218 PeerPool::new(
219 NetworkType::Mainnet,
220 create_generated_tls().expect("generate a TLS identity"),
221 0,
222 requirement,
223 Duration::from_millis(1),
224 )
225 .await
226 }
227
228 #[tokio::test]
230 async fn empty_pool_is_fatal_when_peers_are_required() {
231 assert!(matches!(
232 pool_with_no_connection_attempts(PeerRequirement::Required).await,
233 Err(ChiaQueryError::PeerDiscoveryFailed)
234 ));
235 }
236
237 #[tokio::test]
239 async fn empty_pool_is_tolerated_when_peers_are_optional() {
240 let pool = pool_with_no_connection_attempts(PeerRequirement::Optional)
241 .await
242 .expect("an optional peer pool must construct with zero peers");
243 assert!(!pool.has_peers().await);
244 }
245}