use crate::time::Instant;
use alloy_json_rpc::{RequestPacket, ResponsePacket};
use core::time::Duration;
use derive_more::{Deref, DerefMut};
use futures::{stream::FuturesUnordered, StreamExt};
use parking_lot::RwLock;
use std::{
collections::{HashSet, VecDeque},
num::NonZeroUsize,
sync::Arc,
task::{Context, Poll},
};
use tower::{Layer, Service};
use tracing::trace;
use crate::{TransportError, TransportErrorKind, TransportFut};
const STABILITY_WEIGHT: f64 = 0.7;
const LATENCY_WEIGHT: f64 = 0.3;
const DEFAULT_SAMPLE_COUNT: usize = 10;
const DEFAULT_ACTIVE_TRANSPORT_COUNT: usize = 3;
#[derive(Debug, Clone)]
pub struct FallbackService<S> {
transports: Arc<Vec<ScoredTransport<S>>>,
active_transport_count: usize,
sequential_methods: Arc<HashSet<String>>,
}
impl<S: Clone> FallbackService<S> {
pub fn new(transports: Vec<S>, active_transport_count: usize) -> Self {
Self::new_with_sequential_methods(
transports,
active_transport_count,
default_sequential_methods(),
)
}
pub fn new_with_sequential_methods(
transports: Vec<S>,
active_transport_count: usize,
sequential_methods: HashSet<String>,
) -> Self {
let scored_transports = transports
.into_iter()
.enumerate()
.map(|(id, transport)| ScoredTransport::new(id, transport))
.collect::<Vec<_>>();
Self {
transports: Arc::new(scored_transports),
active_transport_count,
sequential_methods: Arc::new(sequential_methods),
}
}
pub fn append_sequential_method(mut self, sequential_method: impl Into<String>) -> Self {
let mut methods = Arc::unwrap_or_clone(self.sequential_methods);
methods.insert(sequential_method.into());
self.sequential_methods = Arc::new(methods);
self
}
pub fn with_sequential_methods(mut self, sequential_methods: HashSet<String>) -> Self {
self.sequential_methods = Arc::new(sequential_methods);
self
}
fn log_transport_rankings(&self) {
if !tracing::enabled!(tracing::Level::TRACE) {
return;
}
let mut ranked: Vec<(usize, f64, String)> =
self.transports.iter().map(|t| (t.id, t.score(), t.metrics_summary())).collect();
ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
trace!("Current transport rankings:");
for (idx, (id, _score, summary)) in ranked.iter().enumerate() {
trace!(" #{}: Transport[{}] - {}", idx + 1, id, summary);
}
}
fn top_transports(&self) -> Vec<ScoredTransport<S>> {
let mut transports_clone = (*self.transports).clone();
transports_clone.sort_by(|a, b| b.cmp(a));
transports_clone.truncate(self.active_transport_count);
transports_clone
}
}
impl<S> FallbackService<S>
where
S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
+ Send
+ Clone
+ 'static,
{
async fn make_request(&self, req: RequestPacket) -> Result<ResponsePacket, TransportError> {
if req.method_names().any(|name| self.sequential_methods.contains(name)) {
return self.make_request_sequential(req).await;
}
let top_transports = self.top_transports();
if top_transports.is_empty() {
return Err(TransportErrorKind::custom_str(
"No transports available for fallback service",
));
}
let mut futures = FuturesUnordered::new();
for mut transport in top_transports {
let req_clone = req.clone();
let future = async move {
let start = Instant::now();
let result = transport.call(req_clone).await;
trace!(
"Transport[{}] completed: latency={:?}, status={}",
transport.id,
start.elapsed(),
if result.is_ok() { "success" } else { "fail" }
);
(result, transport, start.elapsed())
};
futures.push(future);
}
let mut last_error = None;
while let Some((result, transport, duration)) = futures.next().await {
match result {
Ok(response) => {
transport.track_success(duration);
self.log_transport_rankings();
return Ok(response);
}
Err(error) => {
transport.track_failure();
last_error = Some(error);
}
}
}
Err(last_error.unwrap_or_else(|| {
TransportErrorKind::custom_str("All transport futures failed to complete")
}))
}
async fn make_request_sequential(
&self,
req: RequestPacket,
) -> Result<ResponsePacket, TransportError> {
trace!("Using sequential fallback for method with non-deterministic results");
let top_transports = self.top_transports();
if top_transports.is_empty() {
return Err(TransportErrorKind::custom_str(
"No transports available for fallback service",
));
}
let mut last_error = None;
for mut transport in top_transports {
let req_clone = req.clone();
let start = Instant::now();
trace!("Trying transport[{}] sequentially", transport.id);
match transport.call(req_clone).await {
Ok(response) => {
transport.track_success(start.elapsed());
trace!("Transport[{}] succeeded in {:?}", transport.id, start.elapsed());
self.log_transport_rankings();
return Ok(response);
}
Err(error) => {
transport.track_failure();
trace!("Transport[{}] failed: {:?}, trying next", transport.id, error);
last_error = Some(error);
}
}
}
Err(last_error.unwrap_or_else(|| {
TransportErrorKind::custom_str("All transports failed for sequential request")
}))
}
}
impl<S> Service<RequestPacket> for FallbackService<S>
where
S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
+ Send
+ Sync
+ Clone
+ 'static,
{
type Response = ResponsePacket;
type Error = TransportError;
type Future = TransportFut<'static>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: RequestPacket) -> Self::Future {
let this = self.clone();
Box::pin(async move { this.make_request(req).await })
}
}
#[derive(Debug, Clone)]
pub struct FallbackLayer {
active_transport_count: usize,
sequential_methods: HashSet<String>,
}
impl FallbackLayer {
pub const fn with_active_transport_count(mut self, count: NonZeroUsize) -> Self {
self.active_transport_count = count.get();
self
}
pub fn with_sequential_method(mut self, method: impl Into<String>) -> Self {
self.sequential_methods.insert(method.into());
self
}
pub fn with_sequential_methods(mut self, methods: HashSet<String>) -> Self {
self.sequential_methods = methods;
self
}
pub fn without_sequential_methods(mut self) -> Self {
self.sequential_methods.clear();
self
}
}
impl<S> Layer<Vec<S>> for FallbackLayer
where
S: Service<RequestPacket, Future = TransportFut<'static>, Error = TransportError>
+ Send
+ Clone
+ 'static,
{
type Service = FallbackService<S>;
fn layer(&self, inner: Vec<S>) -> Self::Service {
FallbackService::new_with_sequential_methods(
inner,
self.active_transport_count,
self.sequential_methods.clone(),
)
}
}
impl Default for FallbackLayer {
fn default() -> Self {
Self {
active_transport_count: DEFAULT_ACTIVE_TRANSPORT_COUNT,
sequential_methods: default_sequential_methods(),
}
}
}
#[derive(Debug, Clone, Deref, DerefMut)]
struct ScoredTransport<S> {
#[deref]
#[deref_mut]
transport: S,
id: usize,
metrics: Arc<RwLock<TransportMetrics>>,
}
impl<S> ScoredTransport<S> {
fn new(id: usize, transport: S) -> Self {
Self { id, transport, metrics: Arc::new(Default::default()) }
}
fn score(&self) -> f64 {
let metrics = self.metrics.read();
metrics.calculate_score()
}
fn metrics_summary(&self) -> String {
let metrics = self.metrics.read();
metrics.get_summary()
}
fn track_success(&self, duration: Duration) {
let mut metrics = self.metrics.write();
metrics.track_success(duration);
}
fn track_failure(&self) {
let mut metrics = self.metrics.write();
metrics.track_failure();
}
}
impl<S> PartialEq for ScoredTransport<S> {
fn eq(&self, other: &Self) -> bool {
self.score().eq(&other.score())
}
}
impl<S> Eq for ScoredTransport<S> {}
#[expect(clippy::non_canonical_partial_ord_impl)]
impl<S> PartialOrd for ScoredTransport<S> {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
self.score().partial_cmp(&other.score())
}
}
impl<S> Ord for ScoredTransport<S> {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.partial_cmp(other).unwrap_or(std::cmp::Ordering::Equal)
}
}
#[derive(Debug)]
struct TransportMetrics {
latencies: VecDeque<Duration>,
successes: VecDeque<bool>,
last_update: Instant,
total_requests: u64,
successful_requests: u64,
}
impl TransportMetrics {
fn track_success(&mut self, duration: Duration) {
self.total_requests += 1;
self.successful_requests += 1;
self.last_update = Instant::now();
self.latencies.push_back(duration);
self.successes.push_back(true);
while self.latencies.len() > DEFAULT_SAMPLE_COUNT {
self.latencies.pop_front();
}
while self.successes.len() > DEFAULT_SAMPLE_COUNT {
self.successes.pop_front();
}
}
fn track_failure(&mut self) {
self.total_requests += 1;
self.last_update = Instant::now();
self.successes.push_back(false);
while self.successes.len() > DEFAULT_SAMPLE_COUNT {
self.successes.pop_front();
}
}
fn calculate_score(&self) -> f64 {
if self.successes.is_empty() {
return 0.0;
}
let success_count = self.successes.iter().filter(|&&s| s).count();
let stability_score = success_count as f64 / self.successes.len() as f64;
let latency_score = if !self.latencies.is_empty() {
let avg_latency = self.latencies.iter().map(|d| d.as_secs_f64()).sum::<f64>()
/ self.latencies.len() as f64;
1.0 / (1.0 + avg_latency)
} else {
0.0
};
(stability_score * STABILITY_WEIGHT) + (latency_score * LATENCY_WEIGHT)
}
fn get_summary(&self) -> String {
let success_rate = if !self.successes.is_empty() {
let success_count = self.successes.iter().filter(|&&s| s).count();
success_count as f64 / self.successes.len() as f64
} else {
0.0
};
let avg_latency = if !self.latencies.is_empty() {
self.latencies.iter().map(|d| d.as_secs_f64()).sum::<f64>()
/ self.latencies.len() as f64
} else {
0.0
};
format!(
"success_rate: {:.2}%, avg_latency: {:.2}ms, samples: {}, score: {:.4}",
success_rate * 100.0,
avg_latency * 1000.0,
self.successes.len(),
self.calculate_score()
)
}
}
impl Default for TransportMetrics {
fn default() -> Self {
Self {
latencies: VecDeque::new(),
successes: VecDeque::new(),
last_update: Instant::now(),
total_requests: 0,
successful_requests: 0,
}
}
}
fn default_sequential_methods() -> HashSet<String> {
["eth_sendRawTransactionSync".to_string(), "eth_sendTransactionSync".to_string()]
.into_iter()
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use alloy_json_rpc::{Id, Request, Response, ResponsePayload};
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::time::{sleep, Duration};
use tower::Service;
#[derive(Clone)]
struct DelayedMockTransport {
delay: Duration,
response: Arc<RwLock<Option<ResponsePayload>>>,
call_count: Arc<AtomicUsize>,
}
impl DelayedMockTransport {
fn new(delay: Duration, response: ResponsePayload) -> Self {
Self {
delay,
response: Arc::new(RwLock::new(Some(response))),
call_count: Arc::new(AtomicUsize::new(0)),
}
}
fn call_count(&self) -> usize {
self.call_count.load(Ordering::SeqCst)
}
}
impl Service<RequestPacket> for DelayedMockTransport {
type Response = ResponsePacket;
type Error = TransportError;
type Future = TransportFut<'static>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, req: RequestPacket) -> Self::Future {
self.call_count.fetch_add(1, Ordering::SeqCst);
let delay = self.delay;
let response = self.response.clone();
Box::pin(async move {
sleep(delay).await;
match req {
RequestPacket::Single(single) => {
let resp = response.read().clone().ok_or_else(|| {
TransportErrorKind::custom_str("No response configured")
})?;
Ok(ResponsePacket::Single(Response {
id: single.id().clone(),
payload: resp,
}))
}
RequestPacket::Batch(batch) => {
let resp = response.read().clone().ok_or_else(|| {
TransportErrorKind::custom_str("No response configured")
})?;
let responses = batch
.iter()
.map(|req| Response { id: req.id().clone(), payload: resp.clone() })
.collect();
Ok(ResponsePacket::Batch(responses))
}
}
})
}
}
fn success_response(data: &str) -> ResponsePayload {
let raw = serde_json::value::RawValue::from_string(format!("\"{}\"", data)).unwrap();
ResponsePayload::Success(raw)
}
#[tokio::test]
async fn test_non_deterministic_method_uses_sequential_fallback() {
let transport_a = DelayedMockTransport::new(
Duration::from_millis(50),
success_response("0x1234567890abcdef"), );
let transport_b = DelayedMockTransport::new(
Duration::from_millis(10),
success_response("already_known"), );
let transports = vec![transport_a.clone(), transport_b.clone()];
let mut fallback_service = FallbackService::new(transports, 2);
let request = Request::new(
"eth_sendRawTransactionSync",
Id::Number(1),
[serde_json::Value::String("0xabcdef".to_string())],
);
let serialized = request.serialize().unwrap();
let request_packet = RequestPacket::Single(serialized);
let start = std::time::Instant::now();
let response = fallback_service.call(request_packet).await.unwrap();
let elapsed = start.elapsed();
let result = match response {
ResponsePacket::Single(resp) => match resp.payload {
ResponsePayload::Success(data) => data.get().to_string(),
ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
},
ResponsePacket::Batch(_) => panic!("Unexpected batch response"),
};
assert_eq!(transport_a.call_count(), 1, "First transport should be called");
assert_eq!(transport_b.call_count(), 0, "Second transport should NOT be called");
assert_eq!(result, "\"0x1234567890abcdef\"");
assert!(
elapsed >= Duration::from_millis(40),
"Should wait for first transport: {:?}",
elapsed
);
}
#[tokio::test]
async fn test_deterministic_method_uses_parallel_execution() {
let tx_hash = "0x1234567890abcdef1234567890abcdef1234567890abcdef1234567890abcdef";
let transport_a = DelayedMockTransport::new(
Duration::from_millis(100),
success_response(tx_hash), );
let transport_b = DelayedMockTransport::new(
Duration::from_millis(20),
success_response(tx_hash), );
let transports = vec![transport_a.clone(), transport_b.clone()];
let mut fallback_service = FallbackService::new(transports, 2);
let request = Request::new(
"eth_sendRawTransaction",
Id::Number(1),
[serde_json::Value::String("0xabcdef".to_string())],
);
let serialized = request.serialize().unwrap();
let request_packet = RequestPacket::Single(serialized);
let start = std::time::Instant::now();
let response = fallback_service.call(request_packet).await.unwrap();
let elapsed = start.elapsed();
let result = match response {
ResponsePacket::Single(resp) => match resp.payload {
ResponsePayload::Success(data) => data.get().to_string(),
ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
},
ResponsePacket::Batch(_) => panic!("Unexpected batch response"),
};
assert_eq!(transport_a.call_count(), 1, "Transport A should be called");
assert_eq!(transport_b.call_count(), 1, "Transport B should be called");
assert_eq!(result, format!("\"{}\"", tx_hash));
assert!(
elapsed < Duration::from_millis(50),
"Should use parallel execution and return fast: {:?}",
elapsed
);
}
#[tokio::test]
async fn test_batch_with_any_sequential_method_uses_sequential_execution() {
let tx_hash = "0x1234567890abcdef1234567890abcdef1234567890abcdef1234567890abcdef";
let transport_a =
DelayedMockTransport::new(Duration::from_millis(10), success_response(tx_hash));
let transport_b = DelayedMockTransport::new(
Duration::from_millis(10),
success_response("should_not_be_called"),
);
let transports = vec![transport_a.clone(), transport_b.clone()];
let mut fallback_service = FallbackService::new(transports, 2);
let request1 = Request::new("eth_blockNumber", Id::Number(1), ());
let request2 = Request::new(
"eth_sendRawTransactionSync",
Id::Number(2),
[serde_json::Value::String("0xabcdef".to_string())],
);
let batch = vec![request1.serialize().unwrap(), request2.serialize().unwrap()];
let request_packet = RequestPacket::Batch(batch);
let start = std::time::Instant::now();
let response = fallback_service.call(request_packet).await.unwrap();
let elapsed = start.elapsed();
assert_eq!(
transport_a.call_count(),
1,
"Transport A should be called once (first in sequence)"
);
assert_eq!(
transport_b.call_count(),
0,
"Transport B should NOT be called (transport A succeeded)"
);
match response {
ResponsePacket::Batch(responses) => {
assert_eq!(responses.len(), 2, "Should get 2 responses in batch");
for resp in responses {
match resp.payload {
ResponsePayload::Success(_) => {} ResponsePayload::Failure(err) => panic!("Unexpected error: {:?}", err),
}
}
}
ResponsePacket::Single(_) => panic!("Expected batch response"),
}
assert!(
elapsed < Duration::from_millis(50),
"Sequential execution with fast first transport should be quick: {:?}",
elapsed
);
}
#[tokio::test]
async fn test_custom_sequential_method() {
let transport_a =
DelayedMockTransport::new(Duration::from_millis(10), success_response("result_a"));
let transport_b =
DelayedMockTransport::new(Duration::from_millis(10), success_response("result_b"));
let transports = vec![transport_a.clone(), transport_b.clone()];
let custom_methods = ["my_custom_method".to_string()].into_iter().collect();
let mut fallback_service =
FallbackService::new(transports, 2).with_sequential_methods(custom_methods);
let request = Request::new("my_custom_method", Id::Number(1), ());
let serialized = request.serialize().unwrap();
let request_packet = RequestPacket::Single(serialized);
let start = std::time::Instant::now();
let _response = fallback_service.call(request_packet).await.unwrap();
let elapsed = start.elapsed();
assert_eq!(
transport_a.call_count(),
1,
"Transport A should be called once (sequential, first transport)"
);
assert_eq!(
transport_b.call_count(),
0,
"Transport B should NOT be called (sequential mode, A succeeded)"
);
assert!(
elapsed < Duration::from_millis(50),
"Sequential execution with fast first transport: {:?}",
elapsed
);
}
}