1use crate::time::Instant;
2use alloy_json_rpc::{RequestPacket, ResponsePacket};
3use core::time::Duration;
4use derive_more::{Deref, DerefMut};
5use futures::{stream::FuturesUnordered, StreamExt};
6use parking_lot::RwLock;
7use std::{
8 collections::{HashSet, VecDeque},
9 num::NonZeroUsize,
10 sync::Arc,
11 task::{Context, Poll},
12};
13use tower::{Layer, Service};
14use tracing::trace;
15
16use crate::{TransportError, TransportErrorKind, TransportFut};
17
18const STABILITY_WEIGHT: f64 = 0.7;
20const LATENCY_WEIGHT: f64 = 0.3;
21const DEFAULT_SAMPLE_COUNT: usize = 10;
22const DEFAULT_ACTIVE_TRANSPORT_COUNT: usize = 3;
23
24#[derive(Debug, Clone)]
30pub struct FallbackService<S> {
31 transports: Arc<Vec<ScoredTransport<S>>>,
33 active_transport_count: usize,
35 sequential_methods: Arc<HashSet<String>>,
38}
39
40impl<S: Clone> FallbackService<S> {
41 pub fn new(transports: Vec<S>, active_transport_count: usize) -> Self {
49 Self::new_with_sequential_methods(
50 transports,
51 active_transport_count,
52 default_sequential_methods(),
53 )
54 }
55
56 pub fn new_with_sequential_methods(
64 transports: Vec<S>,
65 active_transport_count: usize,
66 sequential_methods: HashSet<String>,
67 ) -> Self {
68 let scored_transports = transports
69 .into_iter()
70 .enumerate()
71 .map(|(id, transport)| ScoredTransport::new(id, transport))
72 .collect::<Vec<_>>();
73
74 Self {
75 transports: Arc::new(scored_transports),
76 active_transport_count,
77 sequential_methods: Arc::new(sequential_methods),
78 }
79 }
80
81 pub fn append_sequential_method(mut self, sequential_method: impl Into<String>) -> Self {
83 let mut methods = Arc::unwrap_or_clone(self.sequential_methods);
84 methods.insert(sequential_method.into());
85 self.sequential_methods = Arc::new(methods);
86 self
87 }
88
89 pub fn with_sequential_methods(mut self, sequential_methods: HashSet<String>) -> Self {
92 self.sequential_methods = Arc::new(sequential_methods);
93 self
94 }
95
96 fn log_transport_rankings(&self) {
98 if !tracing::enabled!(tracing::Level::TRACE) {
99 return;
100 }
101
102 let mut ranked: Vec<(usize, f64, String)> =
104 self.transports.iter().map(|t| (t.id, t.score(), t.metrics_summary())).collect();
105
106 ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
107
108 trace!("Current transport rankings:");
109 for (idx, (id, _score, summary)) in ranked.iter().enumerate() {
110 trace!(" #{}: Transport[{}] - {}", idx + 1, id, summary);
111 }
112 }
113
114 fn top_transports(&self) -> Vec<ScoredTransport<S>> {
117 let mut transports_clone = (*self.transports).clone();
119 transports_clone.sort_by(|a, b| b.cmp(a));
120 transports_clone.truncate(self.active_transport_count);
121 transports_clone
122 }
123}
124
125impl<S> FallbackService<S>
126where
127 S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
128 + Send
129 + Clone
130 + 'static,
131{
132 async fn make_request(&self, req: RequestPacket) -> Result<ResponsePacket, TransportError> {
150 if req.method_names().any(|name| self.sequential_methods.contains(name)) {
154 return self.make_request_sequential(req).await;
155 }
156
157 let top_transports = self.top_transports();
160
161 if top_transports.is_empty() {
162 return Err(TransportErrorKind::custom_str(
163 "No transports available for fallback service",
164 ));
165 }
166
167 let mut futures = FuturesUnordered::new();
169
170 for mut transport in top_transports {
172 let req_clone = req.clone();
173
174 let future = async move {
175 let start = Instant::now();
176 let result = transport.call(req_clone).await;
177 trace!(
178 "Transport[{}] completed: latency={:?}, status={}",
179 transport.id,
180 start.elapsed(),
181 if result.is_ok() { "success" } else { "fail" }
182 );
183
184 (result, transport, start.elapsed())
185 };
186
187 futures.push(future);
188 }
189
190 let mut last_error = None;
192
193 while let Some((result, transport, duration)) = futures.next().await {
194 match result {
195 Ok(response) => {
196 transport.track_success(duration);
198
199 self.log_transport_rankings();
200
201 return Ok(response);
202 }
203 Err(error) => {
204 transport.track_failure();
206
207 last_error = Some(error);
208 }
209 }
210 }
211
212 Err(last_error.unwrap_or_else(|| {
213 TransportErrorKind::custom_str("All transport futures failed to complete")
214 }))
215 }
216
217 async fn make_request_sequential(
225 &self,
226 req: RequestPacket,
227 ) -> Result<ResponsePacket, TransportError> {
228 trace!("Using sequential fallback for method with non-deterministic results");
229
230 let top_transports = self.top_transports();
232
233 if top_transports.is_empty() {
234 return Err(TransportErrorKind::custom_str(
235 "No transports available for fallback service",
236 ));
237 }
238
239 let mut last_error = None;
240
241 for mut transport in top_transports {
243 let req_clone = req.clone();
244 let start = Instant::now();
245
246 trace!("Trying transport[{}] sequentially", transport.id);
247
248 match transport.call(req_clone).await {
249 Ok(response) => {
250 transport.track_success(start.elapsed());
252 trace!("Transport[{}] succeeded in {:?}", transport.id, start.elapsed());
253 self.log_transport_rankings();
254 return Ok(response);
255 }
256 Err(error) => {
257 transport.track_failure();
259 trace!("Transport[{}] failed: {:?}, trying next", transport.id, error);
260 last_error = Some(error);
261 }
262 }
263 }
264
265 Err(last_error.unwrap_or_else(|| {
267 TransportErrorKind::custom_str("All transports failed for sequential request")
268 }))
269 }
270}
271
272impl<S> Service<RequestPacket> for FallbackService<S>
273where
274 S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
275 + Send
276 + Sync
277 + Clone
278 + 'static,
279{
280 type Response = ResponsePacket;
281 type Error = TransportError;
282 type Future = TransportFut<'static>;
283
284 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
285 Poll::Ready(Ok(()))
287 }
288
289 fn call(&mut self, req: RequestPacket) -> Self::Future {
290 let this = self.clone();
291 Box::pin(async move { this.make_request(req).await })
292 }
293}
294
295#[derive(Debug, Clone)]
314pub struct FallbackLayer {
315 active_transport_count: usize,
317 sequential_methods: HashSet<String>,
320}
321
322impl FallbackLayer {
323 pub const fn with_active_transport_count(mut self, count: NonZeroUsize) -> Self {
325 self.active_transport_count = count.get();
326 self
327 }
328
329 pub fn with_sequential_method(mut self, method: impl Into<String>) -> Self {
334 self.sequential_methods.insert(method.into());
335 self
336 }
337
338 pub fn with_sequential_methods(mut self, methods: HashSet<String>) -> Self {
343 self.sequential_methods = methods;
344 self
345 }
346
347 pub fn without_sequential_methods(mut self) -> Self {
352 self.sequential_methods.clear();
353 self
354 }
355}
356
357impl<S> Layer<Vec<S>> for FallbackLayer
358where
359 S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
360 + Send
361 + Clone
362 + 'static,
363{
364 type Service = FallbackService<S>;
365
366 fn layer(&self, inner: Vec<S>) -> Self::Service {
367 FallbackService::new_with_sequential_methods(
368 inner,
369 self.active_transport_count,
370 self.sequential_methods.clone(),
371 )
372 }
373}
374
375impl Default for FallbackLayer {
376 fn default() -> Self {
377 Self {
378 active_transport_count: DEFAULT_ACTIVE_TRANSPORT_COUNT,
379 sequential_methods: default_sequential_methods(),
380 }
381 }
382}
383
384#[derive(Debug, Clone, Deref, DerefMut)]
397struct ScoredTransport<S> {
398 #[deref]
400 #[deref_mut]
401 transport: S,
402 id: usize,
404 metrics: Arc<RwLock<TransportMetrics>>,
406}
407
408impl<S> ScoredTransport<S> {
409 fn new(id: usize, transport: S) -> Self {
411 Self { id, transport, metrics: Arc::new(Default::default()) }
412 }
413
414 fn score(&self) -> f64 {
416 let metrics = self.metrics.read();
417 metrics.calculate_score()
418 }
419
420 fn metrics_summary(&self) -> String {
422 let metrics = self.metrics.read();
423 metrics.get_summary()
424 }
425
426 fn track_success(&self, duration: Duration) {
428 let mut metrics = self.metrics.write();
429 metrics.track_success(duration);
430 }
431
432 fn track_failure(&self) {
434 let mut metrics = self.metrics.write();
435 metrics.track_failure();
436 }
437}
438
439impl<S> PartialEq for ScoredTransport<S> {
440 fn eq(&self, other: &Self) -> bool {
441 self.score().eq(&other.score())
442 }
443}
444
445impl<S> Eq for ScoredTransport<S> {}
446
447#[expect(clippy::non_canonical_partial_ord_impl)]
448impl<S> PartialOrd for ScoredTransport<S> {
449 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
450 self.score().partial_cmp(&other.score())
451 }
452}
453
454impl<S> Ord for ScoredTransport<S> {
455 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
456 self.partial_cmp(other).unwrap_or(std::cmp::Ordering::Equal)
457 }
458}
459
460#[derive(Debug)]
462struct TransportMetrics {
463 latencies: VecDeque<Duration>,
465 successes: VecDeque<bool>,
467 last_update: Instant,
469 total_requests: u64,
471 successful_requests: u64,
473}
474
475impl TransportMetrics {
476 fn track_success(&mut self, duration: Duration) {
478 self.total_requests += 1;
479 self.successful_requests += 1;
480 self.last_update = Instant::now();
481
482 self.latencies.push_back(duration);
484 self.successes.push_back(true);
485
486 while self.latencies.len() > DEFAULT_SAMPLE_COUNT {
488 self.latencies.pop_front();
489 }
490 while self.successes.len() > DEFAULT_SAMPLE_COUNT {
491 self.successes.pop_front();
492 }
493 }
494
495 fn track_failure(&mut self) {
497 self.total_requests += 1;
498 self.last_update = Instant::now();
499
500 self.successes.push_back(false);
502
503 while self.successes.len() > DEFAULT_SAMPLE_COUNT {
505 self.successes.pop_front();
506 }
507 }
508
509 fn calculate_score(&self) -> f64 {
511 if self.successes.is_empty() {
513 return 0.0;
514 }
515
516 let success_count = self.successes.iter().filter(|&&s| s).count();
518 let stability_score = success_count as f64 / self.successes.len() as f64;
519
520 let latency_score = if !self.latencies.is_empty() {
522 let avg_latency = self.latencies.iter().map(|d| d.as_secs_f64()).sum::<f64>()
523 / self.latencies.len() as f64;
524
525 1.0 / (1.0 + avg_latency)
527 } else {
528 0.0
529 };
530
531 (stability_score * STABILITY_WEIGHT) + (latency_score * LATENCY_WEIGHT)
533 }
534
535 fn get_summary(&self) -> String {
537 let success_rate = if !self.successes.is_empty() {
538 let success_count = self.successes.iter().filter(|&&s| s).count();
539 success_count as f64 / self.successes.len() as f64
540 } else {
541 0.0
542 };
543
544 let avg_latency = if !self.latencies.is_empty() {
545 self.latencies.iter().map(|d| d.as_secs_f64()).sum::<f64>()
546 / self.latencies.len() as f64
547 } else {
548 0.0
549 };
550
551 format!(
552 "success_rate: {:.2}%, avg_latency: {:.2}ms, samples: {}, score: {:.4}",
553 success_rate * 100.0,
554 avg_latency * 1000.0,
555 self.successes.len(),
556 self.calculate_score()
557 )
558 }
559}
560
561impl Default for TransportMetrics {
562 fn default() -> Self {
563 Self {
564 latencies: VecDeque::new(),
565 successes: VecDeque::new(),
566 last_update: Instant::now(),
567 total_requests: 0,
568 successful_requests: 0,
569 }
570 }
571}
572
573fn default_sequential_methods() -> HashSet<String> {
592 ["eth_sendRawTransactionSync".to_string(), "eth_sendTransactionSync".to_string()]
593 .into_iter()
594 .collect()
595}
596
597#[cfg(test)]
598mod tests {
599 use super::*;
600 use alloy_json_rpc::{Id, Request, Response, ResponsePayload};
601 use std::sync::atomic::{AtomicUsize, Ordering};
602 use tokio::time::{sleep, Duration};
603 use tower::Service;
604
605 #[derive(Clone)]
607 struct DelayedMockTransport {
608 delay: Duration,
609 response: Arc<RwLock<Option<ResponsePayload>>>,
610 call_count: Arc<AtomicUsize>,
611 }
612
613 impl DelayedMockTransport {
614 fn new(delay: Duration, response: ResponsePayload) -> Self {
615 Self {
616 delay,
617 response: Arc::new(RwLock::new(Some(response))),
618 call_count: Arc::new(AtomicUsize::new(0)),
619 }
620 }
621
622 fn call_count(&self) -> usize {
623 self.call_count.load(Ordering::SeqCst)
624 }
625 }
626
627 impl Service<RequestPacket> for DelayedMockTransport {
628 type Response = ResponsePacket;
629 type Error = TransportError;
630 type Future = TransportFut<'static>;
631
632 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
633 Poll::Ready(Ok(()))
634 }
635
636 fn call(&mut self, req: RequestPacket) -> Self::Future {
637 self.call_count.fetch_add(1, Ordering::SeqCst);
638 let delay = self.delay;
639 let response = self.response.clone();
640
641 Box::pin(async move {
642 sleep(delay).await;
643
644 match req {
645 RequestPacket::Single(single) => {
646 let resp = response.read().clone().ok_or_else(|| {
647 TransportErrorKind::custom_str("No response configured")
648 })?;
649
650 Ok(ResponsePacket::Single(Response {
651 id: single.id().clone(),
652 payload: resp,
653 }))
654 }
655 RequestPacket::Batch(batch) => {
656 let resp = response.read().clone().ok_or_else(|| {
657 TransportErrorKind::custom_str("No response configured")
658 })?;
659
660 let responses = batch
662 .iter()
663 .map(|req| Response { id: req.id().clone(), payload: resp.clone() })
664 .collect();
665
666 Ok(ResponsePacket::Batch(responses))
667 }
668 }
669 })
670 }
671 }
672
673 fn success_response(data: &str) -> ResponsePayload {
675 let raw = serde_json::value::RawValue::from_string(format!("\"{}\"", data)).unwrap();
676 ResponsePayload::Success(raw)
677 }
678
679 #[tokio::test]
680 async fn test_non_deterministic_method_uses_sequential_fallback() {
681 let transport_a = DelayedMockTransport::new(
686 Duration::from_millis(50),
687 success_response("0x1234567890abcdef"), );
689
690 let transport_b = DelayedMockTransport::new(
691 Duration::from_millis(10),
692 success_response("already_known"), );
694
695 let transports = vec![transport_a.clone(), transport_b.clone()];
696 let mut fallback_service = FallbackService::new(transports, 2);
697
698 let request = Request::new(
699 "eth_sendRawTransactionSync",
700 Id::Number(1),
701 [serde_json::Value::String("0xabcdef".to_string())],
702 );
703 let serialized = request.serialize().unwrap();
704 let request_packet = RequestPacket::Single(serialized);
705
706 let start = std::time::Instant::now();
707 let response = fallback_service.call(request_packet).await.unwrap();
708 let elapsed = start.elapsed();
709
710 let result = match response {
711 ResponsePacket::Single(resp) => match resp.payload {
712 ResponsePayload::Success(data) => data.get().to_string(),
713 ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
714 },
715 ResponsePacket::Batch(_) => panic!("Unexpected batch response"),
716 };
717
718 assert_eq!(transport_a.call_count(), 1, "First transport should be called");
720 assert_eq!(transport_b.call_count(), 0, "Second transport should NOT be called");
722
723 assert_eq!(result, "\"0x1234567890abcdef\"");
725
726 assert!(
728 elapsed >= Duration::from_millis(40),
729 "Should wait for first transport: {:?}",
730 elapsed
731 );
732 }
733
734 #[tokio::test]
735 async fn test_deterministic_method_uses_parallel_execution() {
736 let tx_hash = "0x1234567890abcdef1234567890abcdef1234567890abcdef1234567890abcdef";
740
741 let transport_a = DelayedMockTransport::new(
742 Duration::from_millis(100),
743 success_response(tx_hash), );
745
746 let transport_b = DelayedMockTransport::new(
747 Duration::from_millis(20),
748 success_response(tx_hash), );
750
751 let transports = vec![transport_a.clone(), transport_b.clone()];
752 let mut fallback_service = FallbackService::new(transports, 2);
753
754 let request = Request::new(
755 "eth_sendRawTransaction",
756 Id::Number(1),
757 [serde_json::Value::String("0xabcdef".to_string())],
758 );
759 let serialized = request.serialize().unwrap();
760 let request_packet = RequestPacket::Single(serialized);
761
762 let start = std::time::Instant::now();
763 let response = fallback_service.call(request_packet).await.unwrap();
764 let elapsed = start.elapsed();
765
766 let result = match response {
767 ResponsePacket::Single(resp) => match resp.payload {
768 ResponsePayload::Success(data) => data.get().to_string(),
769 ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
770 },
771 ResponsePacket::Batch(_) => panic!("Unexpected batch response"),
772 };
773
774 assert_eq!(transport_a.call_count(), 1, "Transport A should be called");
776 assert_eq!(transport_b.call_count(), 1, "Transport B should be called");
777
778 assert_eq!(result, format!("\"{}\"", tx_hash));
780
781 assert!(
783 elapsed < Duration::from_millis(50),
784 "Should use parallel execution and return fast: {:?}",
785 elapsed
786 );
787 }
788
789 #[tokio::test]
790 async fn test_batch_with_any_sequential_method_uses_sequential_execution() {
791 let tx_hash = "0x1234567890abcdef1234567890abcdef1234567890abcdef1234567890abcdef";
795
796 let transport_a =
798 DelayedMockTransport::new(Duration::from_millis(10), success_response(tx_hash));
799
800 let transport_b = DelayedMockTransport::new(
803 Duration::from_millis(10),
804 success_response("should_not_be_called"),
805 );
806
807 let transports = vec![transport_a.clone(), transport_b.clone()];
808 let mut fallback_service = FallbackService::new(transports, 2);
809
810 let request1 = Request::new("eth_blockNumber", Id::Number(1), ());
814 let request2 = Request::new(
815 "eth_sendRawTransactionSync",
816 Id::Number(2),
817 [serde_json::Value::String("0xabcdef".to_string())],
818 );
819
820 let batch = vec![request1.serialize().unwrap(), request2.serialize().unwrap()];
821 let request_packet = RequestPacket::Batch(batch);
822
823 let start = std::time::Instant::now();
824 let response = fallback_service.call(request_packet).await.unwrap();
825 let elapsed = start.elapsed();
826
827 assert_eq!(
830 transport_a.call_count(),
831 1,
832 "Transport A should be called once (first in sequence)"
833 );
834 assert_eq!(
835 transport_b.call_count(),
836 0,
837 "Transport B should NOT be called (transport A succeeded)"
838 );
839
840 match response {
842 ResponsePacket::Batch(responses) => {
843 assert_eq!(responses.len(), 2, "Should get 2 responses in batch");
844 for resp in responses {
846 match resp.payload {
847 ResponsePayload::Success(_) => {} ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
849 }
850 }
851 }
852 ResponsePacket::Single(_) => panic!("Expected batch response"),
853 }
854
855 assert!(
857 elapsed < Duration::from_millis(50),
858 "Sequential execution with fast first transport should be quick: {:?}",
859 elapsed
860 );
861 }
862
863 #[tokio::test]
864 async fn test_custom_sequential_method() {
865 let transport_a =
869 DelayedMockTransport::new(Duration::from_millis(10), success_response("result_a"));
870
871 let transport_b =
873 DelayedMockTransport::new(Duration::from_millis(10), success_response("result_b"));
874
875 let transports = vec![transport_a.clone(), transport_b.clone()];
876
877 let custom_methods = ["my_custom_method".to_string()].into_iter().collect();
879 let mut fallback_service =
880 FallbackService::new(transports, 2).with_sequential_methods(custom_methods);
881
882 let request = Request::new("my_custom_method", Id::Number(1), ());
883 let serialized = request.serialize().unwrap();
884 let request_packet = RequestPacket::Single(serialized);
885
886 let start = std::time::Instant::now();
887 let _response = fallback_service.call(request_packet).await.unwrap();
888 let elapsed = start.elapsed();
889
890 assert_eq!(
894 transport_a.call_count(),
895 1,
896 "Transport A should be called once (sequential, first transport)"
897 );
898 assert_eq!(
899 transport_b.call_count(),
900 0,
901 "Transport B should NOT be called (sequential mode, A succeeded)"
902 );
903
904 assert!(
906 elapsed < Duration::from_millis(50),
907 "Sequential execution with fast first transport: {:?}",
908 elapsed
909 );
910 }
911}