use crate::{CoreError as Error, Result};
use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, HashMap};
use std::hash::{Hash, Hasher};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum HashStrategy {
Modulo,
ConsistentHash,
Range,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ShardConfig {
pub shard_count: usize,
pub replication_factor: usize,
pub hash_strategy: HashStrategy,
}
impl Default for ShardConfig {
fn default() -> Self {
Self {
shard_count: 4,
replication_factor: 2,
hash_strategy: HashStrategy::ConsistentHash,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Shard {
pub id: String,
pub connection_string: String,
pub weight: f64,
pub is_available: bool,
pub key_count: usize,
}
impl Shard {
pub fn new(id: String, connection_string: String) -> Self {
Self {
id,
connection_string,
weight: 1.0,
is_available: true,
key_count: 0,
}
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
struct VirtualNode {
shard_id: String,
node_id: usize,
hash: u64,
}
pub struct ShardManager {
config: ShardConfig,
shards: HashMap<String, Shard>,
virtual_nodes: BTreeMap<u64, String>,
virtual_nodes_per_shard: usize,
}
impl ShardManager {
pub fn new(config: ShardConfig) -> Self {
Self {
config,
shards: HashMap::new(),
virtual_nodes: BTreeMap::new(),
virtual_nodes_per_shard: 150, }
}
pub fn add_shard(&mut self, id: &str, connection_string: &str) -> Result<()> {
let shard = Shard::new(id.to_string(), connection_string.to_string());
self.shards.insert(id.to_string(), shard);
if self.config.hash_strategy == HashStrategy::ConsistentHash {
self.add_virtual_nodes(id);
}
Ok(())
}
pub fn remove_shard(&mut self, id: &str) -> Result<()> {
self.shards.remove(id);
if self.config.hash_strategy == HashStrategy::ConsistentHash {
self.remove_virtual_nodes(id);
}
Ok(())
}
pub fn get_shard_for_key(&self, table: &str, key: &dyn std::fmt::Display) -> Result<String> {
if self.shards.is_empty() {
return Err(Error::Validation("No shards available".to_string()));
}
let shard_key = format!("{}:{}", table, key);
match self.config.hash_strategy {
HashStrategy::Modulo => self.get_shard_modulo(&shard_key),
HashStrategy::ConsistentHash => self.get_shard_consistent_hash(&shard_key),
HashStrategy::Range => self.get_shard_range(&shard_key),
}
}
pub fn get_all_shards(&self) -> Vec<String> {
self.shards
.values()
.filter(|s| s.is_available)
.map(|s| s.id.clone())
.collect()
}
pub fn get_shard_stats(&self) -> Vec<ShardStats> {
self.shards
.values()
.map(|shard| ShardStats {
shard_id: shard.id.clone(),
key_count: shard.key_count,
is_available: shard.is_available,
weight: shard.weight,
})
.collect()
}
pub fn plan_rebalance(&self) -> Result<RebalancePlan> {
let stats = self.get_shard_stats();
let total_keys: usize = stats.iter().map(|s| s.key_count).sum();
let avg_keys = if self.shards.is_empty() {
0
} else {
total_keys / self.shards.len()
};
let mut moves = Vec::new();
let overloaded: Vec<_> = stats
.iter()
.filter(|s| s.key_count > avg_keys * 12 / 10) .collect();
let underloaded: Vec<_> = stats
.iter()
.filter(|s| s.key_count < avg_keys * 8 / 10) .collect();
for over in &overloaded {
for under in &underloaded {
let keys_to_move = (over.key_count - avg_keys).min(avg_keys - under.key_count);
if keys_to_move > 0 {
moves.push(RebalanceMove {
from_shard: over.shard_id.clone(),
to_shard: under.shard_id.clone(),
estimated_keys: keys_to_move,
});
}
}
}
Ok(RebalancePlan {
total_keys,
avg_keys_per_shard: avg_keys,
moves,
})
}
fn add_virtual_nodes(&mut self, shard_id: &str) {
for i in 0..self.virtual_nodes_per_shard {
let node_key = format!("{}:vnode:{}", shard_id, i);
let hash = self.hash_string(&node_key);
self.virtual_nodes.insert(hash, shard_id.to_string());
}
}
fn remove_virtual_nodes(&mut self, shard_id: &str) {
self.virtual_nodes.retain(|_, sid| sid != shard_id);
}
fn get_shard_modulo(&self, key: &str) -> Result<String> {
let hash = self.hash_string(key);
let shard_index = (hash % self.shards.len() as u64) as usize;
self.shards
.values()
.nth(shard_index)
.map(|s| s.id.clone())
.ok_or_else(|| Error::Validation("Shard not found".to_string()))
}
fn get_shard_consistent_hash(&self, key: &str) -> Result<String> {
if self.virtual_nodes.is_empty() {
return Err(Error::Validation("No virtual nodes configured".to_string()));
}
let hash = self.hash_string(key);
let shard_id = self
.virtual_nodes
.range(hash..)
.next()
.or_else(|| self.virtual_nodes.iter().next()) .map(|(_, sid)| sid.clone())
.ok_or_else(|| Error::Validation("No shard found".to_string()))?;
Ok(shard_id)
}
fn get_shard_range(&self, key: &str) -> Result<String> {
let hash = self.hash_string(key);
let range_size = u64::MAX / self.shards.len() as u64;
let shard_index = (hash / range_size).min(self.shards.len() as u64 - 1) as usize;
self.shards
.values()
.nth(shard_index)
.map(|s| s.id.clone())
.ok_or_else(|| Error::Validation("Shard not found".to_string()))
}
fn hash_string(&self, s: &str) -> u64 {
use std::collections::hash_map::DefaultHasher;
let mut hasher = DefaultHasher::new();
s.hash(&mut hasher);
hasher.finish()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ShardStats {
pub shard_id: String,
pub key_count: usize,
pub is_available: bool,
pub weight: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RebalancePlan {
pub total_keys: usize,
pub avg_keys_per_shard: usize,
pub moves: Vec<RebalanceMove>,
}
impl RebalancePlan {
pub fn is_needed(&self) -> bool {
!self.moves.is_empty()
}
pub fn total_keys_to_move(&self) -> usize {
self.moves.iter().map(|m| m.estimated_keys).sum()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RebalanceMove {
pub from_shard: String,
pub to_shard: String,
pub estimated_keys: usize,
}
pub struct CrossShardQuery {
pub shard_ids: Vec<String>,
pub query_template: String,
}
impl CrossShardQuery {
pub fn new(shard_ids: Vec<String>, query_template: String) -> Self {
Self {
shard_ids,
query_template,
}
}
pub fn shard_count(&self) -> usize {
self.shard_ids.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_shard_manager_creation() {
let config = ShardConfig::default();
let manager = ShardManager::new(config);
assert_eq!(manager.shards.len(), 0);
}
#[test]
fn test_add_shard() {
let config = ShardConfig::default();
let mut manager = ShardManager::new(config);
assert!(
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.is_ok()
);
assert_eq!(manager.shards.len(), 1);
}
#[test]
fn test_remove_shard() {
let config = ShardConfig::default();
let mut manager = ShardManager::new(config);
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.unwrap();
assert_eq!(manager.shards.len(), 1);
assert!(manager.remove_shard("shard_0").is_ok());
assert_eq!(manager.shards.len(), 0);
}
#[test]
fn test_get_shard_for_key_modulo() {
let config = ShardConfig {
hash_strategy: HashStrategy::Modulo,
..Default::default()
};
let mut manager = ShardManager::new(config);
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.unwrap();
manager
.add_shard("shard_1", "postgresql://localhost/db_1")
.unwrap();
let shard_id = manager.get_shard_for_key("users", &"user_123").unwrap();
assert!(shard_id == "shard_0" || shard_id == "shard_1");
let shard_id2 = manager.get_shard_for_key("users", &"user_123").unwrap();
assert_eq!(shard_id, shard_id2);
}
#[test]
fn test_get_shard_for_key_consistent_hash() {
let config = ShardConfig::default(); let mut manager = ShardManager::new(config);
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.unwrap();
manager
.add_shard("shard_1", "postgresql://localhost/db_1")
.unwrap();
let shard_id = manager.get_shard_for_key("users", &"user_123").unwrap();
assert!(shard_id == "shard_0" || shard_id == "shard_1");
let shard_id2 = manager.get_shard_for_key("users", &"user_123").unwrap();
assert_eq!(shard_id, shard_id2);
}
#[test]
fn test_get_all_shards() {
let config = ShardConfig::default();
let mut manager = ShardManager::new(config);
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.unwrap();
manager
.add_shard("shard_1", "postgresql://localhost/db_1")
.unwrap();
let all_shards = manager.get_all_shards();
assert_eq!(all_shards.len(), 2);
}
#[test]
fn test_cross_shard_query() {
let query = CrossShardQuery::new(
vec!["shard_0".to_string(), "shard_1".to_string()],
"SELECT * FROM users WHERE id = ?".to_string(),
);
assert_eq!(query.shard_count(), 2);
}
#[test]
fn test_rebalance_plan() {
let config = ShardConfig::default();
let mut manager = ShardManager::new(config);
manager
.add_shard("shard_0", "postgresql://localhost/db_0")
.unwrap();
manager
.add_shard("shard_1", "postgresql://localhost/db_1")
.unwrap();
if let Some(shard) = manager.shards.get_mut("shard_0") {
shard.key_count = 1000;
}
if let Some(shard) = manager.shards.get_mut("shard_1") {
shard.key_count = 100;
}
let plan = manager.plan_rebalance().unwrap();
assert!(plan.is_needed());
assert!(plan.total_keys_to_move() > 0);
}
}