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),
}
impl AggregateFunction {
pub fn compute_from_rows(
&self,
shard_rows: &[(u32, Vec<serde_json::Value>)],
) -> Option<AggregateValue> {
let all_rows = || shard_rows.iter().flat_map(|(_, rows)| rows.iter());
match self {
AggregateFunction::Count => {
let total: usize = shard_rows.iter().map(|(_, rows)| rows.len()).sum();
Some(AggregateValue::Count(total as i64))
}
AggregateFunction::Sum(col) => {
let sum: f64 = all_rows().filter_map(|r| extract_numeric(r, col)).sum();
Some(AggregateValue::Sum(sum))
}
AggregateFunction::Avg(col) => {
let mut sum = 0.0f64;
let mut n = 0usize;
for v in all_rows().filter_map(|r| extract_numeric(r, col)) {
sum += v;
n += 1;
}
if n == 0 {
None
} else {
Some(AggregateValue::Avg(sum / n as f64))
}
}
AggregateFunction::Min(col) => {
let mut min: Option<f64> = None;
for v in all_rows().filter_map(|r| extract_numeric(r, col)) {
min = Some(min.map_or(v, |m: f64| m.min(v)));
}
min.map(AggregateValue::Min)
}
AggregateFunction::Max(col) => {
let mut max: Option<f64> = None;
for v in all_rows().filter_map(|r| extract_numeric(r, col)) {
max = Some(max.map_or(v, |m: f64| m.max(v)));
}
max.map(AggregateValue::Max)
}
}
}
}
fn extract_numeric(row: &serde_json::Value, col: &str) -> Option<f64> {
match row.get(col) {
Some(serde_json::Value::Number(n)) => n.as_f64(),
_ => None,
}
}
#[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 shard_rows: Vec<(u32, Vec<serde_json::Value>)>,
}
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> {
self.scatter_query_rows(sql, role, None).await
}
pub async fn scatter_query_rows(
&self,
sql: &str,
role: &str,
agg: Option<&AggregateFunction>,
) -> 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.query_rows(&sql).await {
Ok(rows) => {
let count = rows.len() as u64;
Ok((shard_id, count, rows))
}
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 shard_rows = Vec::new();
let mut failed_shards = Vec::new();
let collect_future = async {
while let Some(result) = futures.next().await {
match result {
Ok((shard_id, count, rows)) => {
shard_row_counts.push((shard_id, count));
shard_rows.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()
));
}
let aggregated = agg.and_then(|f| f.compute_from_rows(&shard_rows));
Ok(ScatterResult {
shard_row_counts,
failed_shards,
aggregated,
shard_rows,
})
}
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))
}
}