use std::sync::Arc;
use std::time::Duration;
use futures::stream::{FuturesUnordered, StreamExt};
use crate::database::sharding::ShardRouter;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PartialFailurePolicy {
Fail,
BestEffort,
}
#[derive(Debug, Clone)]
pub enum AggregateFunction {
Count,
Sum(String),
Avg(String),
Min(String),
Max(String),
}
#[derive(Debug, Clone)]
pub enum AggregateValue {
Count(i64),
Sum(f64),
Avg(f64),
Min(f64),
Max(f64),
}
#[derive(Debug, Clone)]
pub struct ShardError {
pub shard_id: u32,
pub error: String,
}
#[derive(Debug)]
pub struct ScatterResult {
pub shard_row_counts: Vec<(u32, u64)>,
pub failed_shards: Vec<ShardError>,
pub aggregated: Option<AggregateValue>,
}
pub struct ScatterGatherExecutor {
router: Arc<ShardRouter>,
timeout: Duration,
partial_failure: PartialFailurePolicy,
}
impl ScatterGatherExecutor {
pub fn new(router: Arc<ShardRouter>, timeout: Duration, partial_failure: PartialFailurePolicy) -> Self {
Self {
router,
timeout,
partial_failure,
}
}
pub async fn scatter_query(&self, sql: &str, role: &str) -> Result<ScatterResult, String> {
let shards = self.router.all_shards();
let mut futures = FuturesUnordered::new();
for shard_info in shards {
let shard_id = shard_info.shard_id;
if let Some(pool) = self.router.get_pool(shard_id) {
let sql = sql.to_string();
let role = role.to_string();
futures.push(async move {
match pool.get_session(&role).await {
Ok(session) => match session.execute_raw(&sql).await {
Ok(exec_result) => Ok((shard_id, exec_result.rows_affected())),
Err(e) => Err(ShardError {
shard_id,
error: e.to_string(),
}),
},
Err(e) => Err(ShardError {
shard_id,
error: e.to_string(),
}),
}
});
}
}
let mut shard_row_counts = Vec::new();
let mut failed_shards = Vec::new();
let collect_future = async {
while let Some(result) = futures.next().await {
match result {
Ok((shard_id, rows)) => shard_row_counts.push((shard_id, rows)),
Err(err) => failed_shards.push(err),
}
}
};
match tokio::time::timeout(self.timeout, collect_future).await {
Ok(()) => {}
Err(_) => return Err("Scatter-gather query timed out".to_string()),
}
if !failed_shards.is_empty() && self.partial_failure == PartialFailurePolicy::Fail {
return Err(format!(
"Scatter-gather failed: {} shard(s) failed",
failed_shards.len()
));
}
Ok(ScatterResult {
shard_row_counts,
failed_shards,
aggregated: None,
})
}
pub fn aggregate_count(result: &ScatterResult) -> AggregateValue {
let total: u64 = result.shard_row_counts.iter().map(|(_, count)| count).sum();
AggregateValue::Count(total as i64)
}
pub fn aggregate_sum(values: &[f64]) -> AggregateValue {
AggregateValue::Sum(values.iter().sum())
}
pub fn aggregate_avg(values: &[f64]) -> AggregateValue {
if values.is_empty() {
AggregateValue::Avg(0.0)
} else {
AggregateValue::Avg(values.iter().sum::<f64>() / values.len() as f64)
}
}
pub fn aggregate_min(values: &[f64]) -> AggregateValue {
AggregateValue::Min(values.iter().copied().fold(f64::INFINITY, f64::min))
}
pub fn aggregate_max(values: &[f64]) -> AggregateValue {
AggregateValue::Max(values.iter().copied().fold(f64::NEG_INFINITY, f64::max))
}
}