use crate::error::{DbError, Result};
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use sqlx::PgPool;
use std::collections::HashMap;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use uuid::Uuid;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ShardKey {
UserId(Uuid),
TokenId(Uuid),
Custom(String),
Integer(i64),
}
impl ShardKey {
pub fn hash_value(&self) -> u64 {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
self.hash(&mut hasher);
hasher.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ShardingStrategy {
Hash,
Range,
Geographic,
Custom,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ShardInfo {
pub id: u32,
pub name: String,
pub connection_string: String,
pub is_active: bool,
pub weight: u32,
pub region: Option<String>,
pub range_start: Option<u64>,
pub range_end: Option<u64>,
}
pub struct ShardPoolManager {
pools: Arc<RwLock<HashMap<u32, PgPool>>>,
shards: Arc<RwLock<Vec<ShardInfo>>>,
strategy: ShardingStrategy,
}
impl ShardPoolManager {
pub fn new(strategy: ShardingStrategy) -> Self {
Self {
pools: Arc::new(RwLock::new(HashMap::new())),
shards: Arc::new(RwLock::new(Vec::new())),
strategy,
}
}
pub async fn add_shard(&self, shard_info: ShardInfo) -> Result<()> {
let pool = PgPool::connect(&shard_info.connection_string)
.await
.map_err(|e| DbError::Connection(format!("Failed to connect to shard: {}", e)))?;
let mut pools = self.pools.write();
let mut shards = self.shards.write();
pools.insert(shard_info.id, pool);
shards.push(shard_info);
Ok(())
}
pub async fn remove_shard(&self, shard_id: u32) -> Result<()> {
let mut pools = self.pools.write();
let mut shards = self.shards.write();
pools.remove(&shard_id);
shards.retain(|s| s.id != shard_id);
Ok(())
}
pub fn get_shard_id(&self, key: &ShardKey) -> Result<u32> {
let shards = self.shards.read();
if shards.is_empty() {
return Err(DbError::Other("No shards configured".to_string()));
}
match self.strategy {
ShardingStrategy::Hash => {
let hash = key.hash_value();
let active_shards: Vec<_> = shards.iter().filter(|s| s.is_active).collect();
if active_shards.is_empty() {
return Err(DbError::Other("No active shards available".to_string()));
}
let index = (hash % active_shards.len() as u64) as usize;
Ok(active_shards[index].id)
}
ShardingStrategy::Range => {
let hash = key.hash_value();
for shard in shards.iter().filter(|s| s.is_active) {
if let (Some(start), Some(end)) = (shard.range_start, shard.range_end) {
if hash >= start && hash < end {
return Ok(shard.id);
}
}
}
Err(DbError::Other(
"No shard found for key in range".to_string(),
))
}
ShardingStrategy::Geographic | ShardingStrategy::Custom => {
shards
.iter()
.find(|s| s.is_active)
.map(|s| s.id)
.ok_or_else(|| DbError::Other("No active shards available".to_string()))
}
}
}
pub fn get_pool(&self, key: &ShardKey) -> Result<PgPool> {
let shard_id = self.get_shard_id(key)?;
self.get_pool_by_id(shard_id)
}
pub fn get_pool_by_id(&self, shard_id: u32) -> Result<PgPool> {
let pools = self.pools.read();
pools
.get(&shard_id)
.cloned()
.ok_or_else(|| DbError::Other(format!("Shard {} not found", shard_id)))
}
pub fn get_all_active_pools(&self) -> Vec<(u32, PgPool)> {
let pools = self.pools.read();
let shards = self.shards.read();
shards
.iter()
.filter(|s| s.is_active)
.filter_map(|s| pools.get(&s.id).map(|pool| (s.id, pool.clone())))
.collect()
}
pub fn get_shard_info(&self, shard_id: u32) -> Option<ShardInfo> {
let shards = self.shards.read();
shards.iter().find(|s| s.id == shard_id).cloned()
}
pub fn list_shards(&self) -> Vec<ShardInfo> {
let shards = self.shards.read();
shards.clone()
}
pub fn set_shard_active(&self, shard_id: u32, active: bool) -> Result<()> {
let mut shards = self.shards.write();
if let Some(shard) = shards.iter_mut().find(|s| s.id == shard_id) {
shard.is_active = active;
Ok(())
} else {
Err(DbError::Other(format!("Shard {} not found", shard_id)))
}
}
pub fn shard_count(&self) -> usize {
let shards = self.shards.read();
shards.len()
}
pub fn active_shard_count(&self) -> usize {
let shards = self.shards.read();
shards.iter().filter(|s| s.is_active).count()
}
}
pub struct ShardCoordinator {
manager: Arc<ShardPoolManager>,
}
impl ShardCoordinator {
pub fn new(manager: Arc<ShardPoolManager>) -> Self {
Self { manager }
}
pub async fn execute_on_shard<F, T>(&self, key: &ShardKey, f: F) -> Result<T>
where
F: FnOnce(
&PgPool,
)
-> std::pin::Pin<Box<dyn std::future::Future<Output = Result<T>> + Send>>
+ Send,
T: Send,
{
let pool = self.manager.get_pool(key)?;
f(&pool).await
}
pub async fn execute_on_all_shards<F, T>(&self, f: F) -> Result<Vec<(u32, T)>>
where
F: Fn(&PgPool) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<T>> + Send>>
+ Send
+ Sync,
T: Send,
{
let pools = self.manager.get_all_active_pools();
let mut results = Vec::new();
for (shard_id, pool) in pools {
match f(&pool).await {
Ok(result) => results.push((shard_id, result)),
Err(e) => {
tracing::warn!("Error executing on shard {}: {}", shard_id, e);
}
}
}
Ok(results)
}
pub async fn aggregate_from_all_shards<F, T, R>(
&self,
query: F,
aggregator: fn(Vec<T>) -> R,
) -> Result<R>
where
F: Fn(&PgPool) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<T>> + Send>>
+ Send
+ Sync,
T: Send,
R: Send,
{
let results = self.execute_on_all_shards(query).await?;
let values: Vec<T> = results.into_iter().map(|(_, v)| v).collect();
Ok(aggregator(values))
}
}
pub struct ConsistentHashRing {
virtual_nodes: u32,
ring: Arc<RwLock<Vec<(u64, u32)>>>,
}
impl ConsistentHashRing {
pub fn new(virtual_nodes: u32) -> Self {
Self {
virtual_nodes,
ring: Arc::new(RwLock::new(Vec::new())),
}
}
pub fn add_shard(&self, shard_id: u32) {
let mut ring = self.ring.write();
for i in 0..self.virtual_nodes {
let key = format!("shard-{}-vnode-{}", shard_id, i);
let mut hasher = std::collections::hash_map::DefaultHasher::new();
key.hash(&mut hasher);
let hash = hasher.finish();
ring.push((hash, shard_id));
}
ring.sort_by_key(|(hash, _)| *hash);
}
pub fn remove_shard(&self, shard_id: u32) {
let mut ring = self.ring.write();
ring.retain(|(_, id)| *id != shard_id);
}
pub fn get_shard(&self, key: &ShardKey) -> Option<u32> {
let ring = self.ring.read();
if ring.is_empty() {
return None;
}
let hash = key.hash_value();
match ring.binary_search_by_key(&hash, |(h, _)| *h) {
Ok(idx) => Some(ring[idx].1),
Err(idx) => {
if idx >= ring.len() {
Some(ring[0].1)
} else {
Some(ring[idx].1)
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_shard_key_hash() {
let key1 = ShardKey::UserId(Uuid::nil());
let key2 = ShardKey::UserId(Uuid::nil());
assert_eq!(key1.hash_value(), key2.hash_value());
}
#[test]
fn test_consistent_hash_ring() {
let ring = ConsistentHashRing::new(100);
ring.add_shard(1);
ring.add_shard(2);
ring.add_shard(3);
let key = ShardKey::UserId(Uuid::nil());
let shard = ring.get_shard(&key);
assert!(shard.is_some());
assert!(shard.unwrap() <= 3);
}
#[test]
fn test_shard_manager_hash_strategy() {
let manager = ShardPoolManager::new(ShardingStrategy::Hash);
assert_eq!(manager.shard_count(), 0);
assert_eq!(manager.active_shard_count(), 0);
}
}