use crate::{ShardingError, ShardingRouter};
use std::collections::HashMap;
use std::sync::Arc;
pub struct ScatterGather {
router: Arc<ShardingRouter>,
}
impl ScatterGather {
pub fn new(router: Arc<ShardingRouter>) -> Self {
Self { router }
}
pub fn broadcast<F, T>(&self, f: F) -> Vec<Result<T, ShardingError>>
where
F: Fn(&str) -> Result<T, ShardingError> + Send + Sync,
T: Send,
{
let shards = self.router.query_all();
if shards.is_empty() {
return Vec::new();
}
std::thread::scope(|s| {
let handles: Vec<_> = shards
.iter()
.map(|shard| s.spawn(|| f(shard.as_str())))
.collect();
handles
.into_iter()
.map(|h| h.join().unwrap_or_else(|_| Err(ShardingError::ThreadPanic)))
.collect()
})
}
pub fn scatter_by_keys<F, T>(&self, keys: &[String], f: F) -> Vec<Result<T, ShardingError>>
where
F: Fn(&str, &[String]) -> Result<T, ShardingError> + Send + Sync,
T: Send,
{
let mut groups: HashMap<String, Vec<String>> = HashMap::new();
for key in keys {
if let Ok(shard) = self.router.route(key) {
groups
.entry(shard.to_string())
.or_default()
.push(key.clone());
}
}
if groups.is_empty() {
return Vec::new();
}
let groups_vec: Vec<(String, Vec<String>)> = groups.into_iter().collect();
std::thread::scope(|s| {
let handles: Vec<_> = groups_vec
.iter()
.map(|(shard, shard_keys)| s.spawn(|| f(shard.as_str(), shard_keys)))
.collect();
handles
.into_iter()
.map(|h| h.join().unwrap_or_else(|_| Err(ShardingError::ThreadPanic)))
.collect()
})
}
pub fn merge<T, F>(
results: Vec<Result<T, ShardingError>>,
merger: F,
) -> Result<T, ShardingError>
where
F: FnOnce(Vec<T>) -> Result<T, ShardingError>,
{
let mut oks: Vec<T> = Vec::with_capacity(results.len());
for r in results {
let t = r?;
oks.push(t);
}
merger(oks)
}
}
impl std::fmt::Debug for ScatterGather {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ScatterGather")
.field("shard_count", &self.router.shard_count())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ShardingStrategy;
use std::collections::HashSet;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::{Duration, Instant};
#[test]
fn test_broadcast_all_shards() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["s0", "s1", "s2"],
));
let sg = ScatterGather::new(router);
let results = sg.broadcast(|shard| Ok(shard.to_string()));
assert_eq!(results.len(), 3);
let mut shards: HashSet<String> = results.into_iter().map(|r| r.unwrap()).collect();
assert!(shards.remove("s0"));
assert!(shards.remove("s1"));
assert!(shards.remove("s2"));
}
#[test]
fn test_broadcast_empty_shards() {
let router = Arc::new(ShardingRouter::new(ShardingStrategy::Hash, vec![]));
let sg = ScatterGather::new(router);
let results: Vec<Result<String, ShardingError>> = sg.broadcast(|_| Ok("x".to_string()));
assert!(results.is_empty());
}
#[test]
fn test_broadcast_single_shard() {
let router = Arc::new(ShardingRouter::new(ShardingStrategy::Hash, vec!["only"]));
let sg = ScatterGather::new(router);
let results = sg.broadcast(|s| Ok(s.to_string()));
assert_eq!(results.len(), 1);
assert_eq!(results[0].as_ref().unwrap(), &"only");
}
#[test]
fn test_broadcast_parallel_execution() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["s0", "s1", "s2"],
));
let sg = ScatterGather::new(router);
let start = Instant::now();
let results = sg.broadcast(|_| {
std::thread::sleep(Duration::from_millis(60));
Ok(1u32)
});
let elapsed = start.elapsed();
assert_eq!(results.len(), 3);
assert!(
elapsed < Duration::from_millis(170),
"并行执行总耗时应远小于串行 180ms,实际: {:?}",
elapsed
);
}
#[test]
fn test_broadcast_collects_values() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["a", "b", "c", "d"],
));
let sg = ScatterGather::new(router);
let results = sg.broadcast(|shard| Ok(shard.len() as i32));
let sum: i32 = results.into_iter().map(|r| r.unwrap()).sum();
assert_eq!(sum, 4);
}
#[test]
fn test_scatter_by_keys_groups() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["s0", "s1", "s2", "s3"],
));
let sg = ScatterGather::new(router);
let keys: Vec<String> = (0..100).map(|i| format!("key_{}", i)).collect();
let results = sg.scatter_by_keys(&keys, |shard, ks| Ok((shard.to_string(), ks.len())));
let mut total = 0usize;
for r in results {
let (shard, count) = r.unwrap();
assert!(shard.starts_with('s'));
total += count;
}
assert_eq!(total, 100);
}
#[test]
fn test_scatter_by_keys_one_call_per_shard() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["s0", "s1"],
));
let sg = ScatterGather::new(router);
let call_count = Arc::new(AtomicUsize::new(0));
let keys: Vec<String> = (0..50).map(|i| format!("k{}", i)).collect();
let cc = call_count.clone();
let results = sg.scatter_by_keys(&keys, move |_shard, ks| {
cc.fetch_add(1, Ordering::SeqCst);
Ok(ks.len())
});
let calls = call_count.load(Ordering::SeqCst);
assert!(
calls == 1 || calls == 2,
"calls should be 1 or 2, got {}",
calls
);
assert_eq!(results.len(), calls);
let total: usize = results.into_iter().map(|r| r.unwrap()).sum();
assert_eq!(total, 50);
}
#[test]
fn test_scatter_by_keys_empty() {
let router = Arc::new(ShardingRouter::new(
ShardingStrategy::Hash,
vec!["s0", "s1"],
));
let sg = ScatterGather::new(router);
let results: Vec<Result<u32, ShardingError>> = sg.scatter_by_keys(&[], |_, _| Ok(1));
assert!(results.is_empty());
}
#[test]
fn test_scatter_by_keys_returns_keys_per_shard() {
let mut keys = HashSet::new();
keys.insert("h1".to_string());
keys.insert("h2".to_string());
let router = Arc::new(ShardingRouter::new_list(
keys,
"hit_shard".to_string(),
Some("default_shard".to_string()),
));
let sg = ScatterGather::new(router);
let input = vec![
"h1".to_string(),
"h2".to_string(),
"miss1".to_string(),
"miss2".to_string(),
];
let results = sg.scatter_by_keys(&input, |shard, ks| Ok((shard.to_string(), ks.to_vec())));
let mut hit_count = 0;
let mut default_count = 0;
for r in results {
let (shard, ks) = r.unwrap();
match shard.as_str() {
"hit_shard" => {
assert_eq!(ks.len(), 2);
assert!(ks.contains(&"h1".to_string()));
assert!(ks.contains(&"h2".to_string()));
hit_count += 1;
}
"default_shard" => {
assert_eq!(ks.len(), 2);
assert!(ks.contains(&"miss1".to_string()));
assert!(ks.contains(&"miss2".to_string()));
default_count += 1;
}
_ => panic!("unexpected shard: {}", shard),
}
}
assert_eq!(hit_count, 1);
assert_eq!(default_count, 1);
}
#[test]
fn test_merge_all_ok() {
let results: Vec<Result<i32, ShardingError>> = vec![Ok(1), Ok(2), Ok(3)];
let merged = ScatterGather::merge(results, |vs| Ok(vs.iter().sum())).unwrap();
assert_eq!(merged, 6);
}
#[test]
fn test_merge_propagates_first_error() {
let results: Vec<Result<i32, ShardingError>> =
vec![Ok(1), Err(ShardingError::ThreadPanic), Ok(3)];
let result = ScatterGather::merge(results, |vs| Ok(vs.iter().sum()));
assert!(matches!(result, Err(ShardingError::ThreadPanic)));
}
#[test]
fn test_merge_empty() {
let results: Vec<Result<i32, ShardingError>> = vec![];
let merged = ScatterGather::merge(results, |vs| Ok(vs.len() as i32)).unwrap();
assert_eq!(merged, 0);
}
#[test]
fn test_merge_single_value() {
let results: Vec<Result<String, ShardingError>> = vec![Ok("only".to_string())];
let merged = ScatterGather::merge(results, |mut vs| Ok(vs.remove(0))).unwrap();
assert_eq!(merged, "only");
}
#[test]
fn test_debug_format() {
let router = Arc::new(ShardingRouter::new(ShardingStrategy::Hash, vec!["s0"]));
let sg = ScatterGather::new(router);
let s = format!("{:?}", sg);
assert!(s.contains("ScatterGather"));
assert!(s.contains("shard_count"));
}
}