1use std::sync::Arc;
10use std::time::Duration;
11
12use async_trait::async_trait;
13use chia_protocol::{Bytes32, CoinState, CoinStateFilters, Program, SpendBundle};
14use chia_wallet_sdk::client::Peer;
15use tokio::sync::RwLock;
16
17use crate::error::ChiaPeerError;
18
19const MAX_PUZZLE_STATE_PAGES: usize = 10_000;
22
23const MAX_ACCUMULATED_COIN_STATES: usize = 500_000;
26
27#[async_trait]
32pub trait CoinStateFetcher: Send + Sync {
33 async fn coin_states(
35 &self,
36 coin_ids: Vec<Bytes32>,
37 subscribe: bool,
38 ) -> Result<Vec<CoinState>, ChiaPeerError>;
39
40 async fn puzzle_states(
43 &self,
44 puzzle_hashes: Vec<Bytes32>,
45 filters: CoinStateFilters,
46 subscribe: bool,
47 ) -> Result<Vec<CoinState>, ChiaPeerError>;
48
49 async fn children(&self, coin_id: Bytes32) -> Result<Vec<CoinState>, ChiaPeerError>;
51
52 async fn puzzle_and_solution(
58 &self,
59 coin_id: Bytes32,
60 height: u32,
61 ) -> Result<(Program, Program), ChiaPeerError>;
62}
63
64#[derive(Clone)]
69pub struct PeerFetcher {
70 peer: Arc<RwLock<Option<Peer>>>,
71 genesis_challenge: Bytes32,
72 request_timeout: Duration,
73}
74
75impl PeerFetcher {
76 pub fn new(peer: Peer, genesis_challenge: Bytes32, request_timeout: Duration) -> Self {
79 Self {
80 peer: Arc::new(RwLock::new(Some(peer))),
81 genesis_challenge,
82 request_timeout,
83 }
84 }
85
86 #[cfg(test)]
88 fn disconnected(genesis_challenge: Bytes32, request_timeout: Duration) -> Self {
89 Self {
90 peer: Arc::new(RwLock::new(None)),
91 genesis_challenge,
92 request_timeout,
93 }
94 }
95
96 pub async fn swap_peer(&self, peer: Peer) {
98 *self.peer.write().await = Some(peer);
99 }
100
101 async fn peer(&self) -> Result<Peer, ChiaPeerError> {
103 self.peer
104 .read()
105 .await
106 .clone()
107 .ok_or(ChiaPeerError::NotConnected)
108 }
109
110 pub async fn send_transaction(&self, bundle: SpendBundle) -> Result<u8, ChiaPeerError> {
112 let peer = self.peer().await?;
113 let ack = self
114 .with_timeout(peer.send_transaction(bundle))
115 .await?
116 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?;
117 Ok(ack.status)
118 }
119
120 pub async fn remove_coin_subscriptions(
122 &self,
123 coin_ids: Vec<Bytes32>,
124 ) -> Result<(), ChiaPeerError> {
125 let peer = self.peer().await?;
126 self.with_timeout(peer.remove_coin_subscriptions(Some(coin_ids)))
127 .await?
128 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?;
129 Ok(())
130 }
131
132 async fn with_timeout<T>(
135 &self,
136 fut: impl std::future::Future<Output = T>,
137 ) -> Result<T, ChiaPeerError> {
138 tokio::time::timeout(self.request_timeout, fut)
139 .await
140 .map_err(|_| ChiaPeerError::Timeout)
141 }
142}
143
144struct PuzzleStatePage {
146 coin_states: Vec<CoinState>,
147 height: u32,
148 header_hash: Bytes32,
149 is_finished: bool,
150}
151
152async fn collect_paged<F, Fut>(
159 genesis_challenge: Bytes32,
160 mut fetch_page: F,
161) -> Result<Vec<CoinState>, ChiaPeerError>
162where
163 F: FnMut(Option<u32>, Bytes32) -> Fut,
164 Fut: std::future::Future<Output = Result<PuzzleStatePage, ChiaPeerError>>,
165{
166 let mut all = Vec::new();
167 let mut previous_height: Option<u32> = None;
168 let mut header_hash = genesis_challenge;
169
170 for _page in 0..MAX_PUZZLE_STATE_PAGES {
171 let page = fetch_page(previous_height, header_hash).await?;
172 all.extend(page.coin_states);
173 if all.len() > MAX_ACCUMULATED_COIN_STATES {
174 return Err(ChiaPeerError::Rejected(format!(
175 "puzzle-state response exceeded {MAX_ACCUMULATED_COIN_STATES} coins"
176 )));
177 }
178 if page.is_finished {
179 return Ok(all);
180 }
181 if previous_height.is_some_and(|prev| page.height <= prev) {
182 return Err(ChiaPeerError::Rejected(
183 "puzzle-state paging did not advance the height".into(),
184 ));
185 }
186 previous_height = Some(page.height);
187 header_hash = page.header_hash;
188 }
189 Err(ChiaPeerError::Rejected(format!(
190 "puzzle-state paging exceeded {MAX_PUZZLE_STATE_PAGES} pages"
191 )))
192}
193
194#[async_trait]
195impl CoinStateFetcher for PeerFetcher {
196 async fn coin_states(
197 &self,
198 coin_ids: Vec<Bytes32>,
199 subscribe: bool,
200 ) -> Result<Vec<CoinState>, ChiaPeerError> {
201 let peer = self.peer().await?;
202 let response = self
203 .with_timeout(peer.request_coin_state(
204 coin_ids,
205 None,
206 self.genesis_challenge,
207 subscribe,
208 ))
209 .await?
210 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?
211 .map_err(|_| ChiaPeerError::Rejected("coin-state request rejected".into()))?;
212 Ok(response.coin_states)
213 }
214
215 async fn puzzle_states(
216 &self,
217 puzzle_hashes: Vec<Bytes32>,
218 filters: CoinStateFilters,
219 subscribe: bool,
220 ) -> Result<Vec<CoinState>, ChiaPeerError> {
221 let peer = self.peer().await?;
222 collect_paged(self.genesis_challenge, |previous_height, header_hash| {
224 let peer = peer.clone();
225 let puzzle_hashes = puzzle_hashes.clone();
226 let filters = filters.clone();
227 let this = self;
228 async move {
229 let response = this
230 .with_timeout(peer.request_puzzle_state(
231 puzzle_hashes,
232 previous_height,
233 header_hash,
234 filters,
235 subscribe,
236 ))
237 .await?
238 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?
239 .map_err(|_| ChiaPeerError::Rejected("puzzle-state request rejected".into()))?;
240 Ok(PuzzleStatePage {
241 coin_states: response.coin_states,
242 height: response.height,
243 header_hash: response.header_hash,
244 is_finished: response.is_finished,
245 })
246 }
247 })
248 .await
249 }
250
251 async fn children(&self, coin_id: Bytes32) -> Result<Vec<CoinState>, ChiaPeerError> {
252 let peer = self.peer().await?;
253 let response = self
254 .with_timeout(peer.request_children(coin_id))
255 .await?
256 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?;
257 Ok(response.coin_states)
258 }
259
260 async fn puzzle_and_solution(
261 &self,
262 coin_id: Bytes32,
263 height: u32,
264 ) -> Result<(Program, Program), ChiaPeerError> {
265 let peer = self.peer().await?;
266 let outcome = self
267 .with_timeout(peer.request_puzzle_and_solution(coin_id, height))
268 .await?
269 .map_err(|e| ChiaPeerError::Transport(e.to_string()))?;
270 match outcome {
271 Ok(response) => Ok((response.puzzle, response.solution)),
272 Err(_) => Err(ChiaPeerError::Rejected(
276 "peer rejected puzzle/solution for a known-spent coin".into(),
277 )),
278 }
279 }
280}
281
282#[cfg(test)]
283mod tests {
284 use super::*;
285 use crate::config::ChiaNetwork;
286 use chia_protocol::{Coin, SpendBundle};
287 use chia_wallet_sdk::test::PeerSimulator;
288 use std::time::Duration;
289
290 fn genesis() -> Bytes32 {
291 ChiaNetwork::Testnet11.genesis_challenge()
292 }
293
294 async fn fetcher_over_sim() -> (PeerSimulator, PeerFetcher, Coin) {
295 let sim = PeerSimulator::new().await.expect("start simulator");
296 let coin = sim.lock().await.new_coin(Bytes32::new([7; 32]), 1_000);
297 let (peer, _receiver) = sim.connect_raw().await.expect("connect to simulator");
298 let fetcher = PeerFetcher::new(peer, genesis(), Duration::from_secs(5));
299 (sim, fetcher, coin)
300 }
301
302 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
303 async fn coin_states_reads_an_inserted_coin() {
304 let (_sim, fetcher, coin) = fetcher_over_sim().await;
305 let states = fetcher
306 .coin_states(vec![coin.coin_id()], false)
307 .await
308 .unwrap();
309 assert_eq!(states.len(), 1);
310 assert_eq!(states[0].coin, coin);
311 }
312
313 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
314 async fn coin_states_of_unknown_coin_is_empty() {
315 let (_sim, fetcher, _coin) = fetcher_over_sim().await;
316 let states = fetcher
317 .coin_states(vec![Bytes32::new([0xee; 32])], false)
318 .await
319 .unwrap();
320 assert!(states.is_empty());
321 }
322
323 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
324 async fn puzzle_states_reads_by_puzzle_hash() {
325 let (_sim, fetcher, coin) = fetcher_over_sim().await;
326 let filters = CoinStateFilters {
327 include_spent: true,
328 include_unspent: true,
329 include_hinted: true,
330 min_amount: 0,
331 };
332 let states = fetcher
333 .puzzle_states(vec![coin.puzzle_hash], filters, false)
334 .await
335 .unwrap();
336 assert!(states.iter().any(|s| s.coin == coin));
337 }
338
339 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
340 async fn children_reads_child_coins_of_a_parent() {
341 let sim = PeerSimulator::new().await.unwrap();
342 let parent = Bytes32::new([9; 32]);
343 let child = Coin::new(parent, Bytes32::new([3; 32]), 5);
344 sim.lock().await.insert_coin(child);
345 let (peer, _receiver) = sim.connect_raw().await.unwrap();
346 let fetcher = PeerFetcher::new(peer, genesis(), Duration::from_secs(5));
347
348 let kids = fetcher.children(parent).await.unwrap();
349 assert_eq!(kids.len(), 1);
350 assert_eq!(kids[0].coin, child);
351 }
352
353 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
354 async fn puzzle_and_solution_of_unspent_coin_never_yields_a_spend() {
355 let (_sim, fetcher, coin) = fetcher_over_sim().await;
356 let result = fetcher.puzzle_and_solution(coin.coin_id(), 1).await;
357 assert!(
360 result.is_err(),
361 "unspent coin must fail closed, not yield a reveal: {result:?}"
362 );
363 }
364
365 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
366 async fn submitting_an_invalid_bundle_returns_a_failure_ack() {
367 let (_sim, fetcher, _coin) = fetcher_over_sim().await;
368 let bundle = SpendBundle::new(vec![], chia::bls::Signature::default());
369 let status = fetcher.send_transaction(bundle).await.unwrap();
370 assert_eq!(status, 3, "an empty bundle is rejected with a failure ack");
371 }
372
373 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
374 async fn subscribe_then_remove_coin_subscriptions_succeeds() {
375 let (_sim, fetcher, coin) = fetcher_over_sim().await;
376 fetcher
377 .coin_states(vec![coin.coin_id()], true)
378 .await
379 .unwrap();
380 fetcher
381 .remove_coin_subscriptions(vec![coin.coin_id()])
382 .await
383 .unwrap();
384 }
385
386 fn page(states: usize, height: u32, is_finished: bool) -> PuzzleStatePage {
387 PuzzleStatePage {
388 coin_states: (0..states)
389 .map(|_| CoinState {
390 coin: Coin::new(Bytes32::new([1; 32]), Bytes32::new([2; 32]), 1),
391 created_height: Some(height),
392 spent_height: None,
393 })
394 .collect(),
395 height,
396 header_hash: Bytes32::new([height as u8; 32]),
397 is_finished,
398 }
399 }
400
401 #[tokio::test]
402 async fn paging_that_never_finishes_fails_closed_not_hangs() {
403 let mut next_height = 0u32;
405 let result = collect_paged(Bytes32::default(), |_prev, _hdr| {
406 next_height += 1;
407 let h = next_height;
408 async move { Ok(page(1, h, false)) }
409 })
410 .await;
411 assert!(
412 matches!(result, Err(ChiaPeerError::Rejected(_))),
413 "{result:?}"
414 );
415 }
416
417 #[tokio::test]
418 async fn paging_without_progress_fails_closed() {
419 let result = collect_paged(Bytes32::default(), |_prev, _hdr| async move {
421 Ok(page(1, 42, false))
422 })
423 .await;
424 assert!(
425 matches!(result, Err(ChiaPeerError::Rejected(_))),
426 "{result:?}"
427 );
428 }
429
430 #[tokio::test]
431 async fn paging_over_coin_cap_fails_closed() {
432 let result = collect_paged(Bytes32::default(), |_prev, _hdr| async move {
433 Ok(page(MAX_ACCUMULATED_COIN_STATES + 1, 1, false))
434 })
435 .await;
436 assert!(
437 matches!(result, Err(ChiaPeerError::Rejected(_))),
438 "{result:?}"
439 );
440 }
441
442 #[tokio::test]
443 async fn paging_finishes_normally_returns_all() {
444 let mut calls = 0u32;
445 let result = collect_paged(Bytes32::default(), |_prev, _hdr| {
446 calls += 1;
447 let finished = calls == 2;
448 let h = calls;
449 async move { Ok(page(1, h, finished)) }
450 })
451 .await
452 .unwrap();
453 assert_eq!(result.len(), 2, "both pages accumulated then finished");
454 }
455
456 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
457 async fn disconnected_fetcher_fails_closed_on_every_read() {
458 let fetcher = PeerFetcher::disconnected(genesis(), Duration::from_secs(1));
459 let id = Bytes32::new([1; 32]);
460 assert_eq!(
461 fetcher.coin_states(vec![id], false).await,
462 Err(ChiaPeerError::NotConnected)
463 );
464 let filters = CoinStateFilters {
465 include_spent: true,
466 include_unspent: true,
467 include_hinted: true,
468 min_amount: 0,
469 };
470 assert_eq!(
471 fetcher.puzzle_states(vec![id], filters, false).await,
472 Err(ChiaPeerError::NotConnected)
473 );
474 assert_eq!(fetcher.children(id).await, Err(ChiaPeerError::NotConnected));
475 assert_eq!(
476 fetcher.puzzle_and_solution(id, 1).await,
477 Err(ChiaPeerError::NotConnected)
478 );
479 assert_eq!(
480 fetcher.remove_coin_subscriptions(vec![id]).await,
481 Err(ChiaPeerError::NotConnected)
482 );
483 assert_eq!(
484 fetcher
485 .send_transaction(SpendBundle::new(vec![], chia::bls::Signature::default()))
486 .await,
487 Err(ChiaPeerError::NotConnected)
488 );
489 }
490}