use std::time::Duration;
use serde::{Deserialize, Serialize};
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn sharding_quickstart_example() {
use datum_agent::{
ClusterAgent, ClusterAgentConfig,
dcp::{DcpJobFactories, DcpServerConfig, DcpTcpServerConfig},
};
use datum_cluster::ClusterConfig;
use datum_cluster_sharding::{
EntityContext, ReplyPort, Sharding, ShardingConfig, ShardingResult,
};
#[derive(Serialize, Deserialize)]
enum CounterMsg {
Add(u64),
Get(ReplyPort<u64>),
}
let agent = ClusterAgent::start(
ClusterAgentConfig {
cluster: ClusterConfig::new("orders-1"),
dcp: DcpServerConfig {
tcp: Some(DcpTcpServerConfig {
addr: "127.0.0.1:0".parse().expect("loopback addr"),
}),
..DcpServerConfig::default()
},
..ClusterAgentConfig::default()
},
DcpJobFactories::new(),
)
.await
.expect("cluster agent starts");
let sharding = Sharding::init(
&agent,
ShardingConfig {
num_shards: 64,
..ShardingConfig::default()
},
)
.expect("sharding region");
sharding
.register_entity_type("counter", |_spawn: EntityContext| {
let mut total = 0_u64;
move |_context: &EntityContext, message: CounterMsg| -> ShardingResult<()> {
match message {
CounterMsg::Add(amount) => total += amount,
CounterMsg::Get(reply) => {
let _ = reply.send(total);
}
}
Ok(())
}
})
.await
.expect("entity type registered");
while !agent
.cluster()
.current_state()
.is_placement_coordinator(agent.cluster().node_id())
{
tokio::time::sleep(Duration::from_millis(10)).await;
}
let cart = sharding.entity_ref::<CounterMsg>("counter", "cart-42");
cart.tell(CounterMsg::Add(3)).await.expect("tell add");
let total = cart
.ask(Duration::from_secs(5), CounterMsg::Get)
.await
.expect("ask get");
assert_eq!(total, 3);
sharding.shutdown().await;
agent.shutdown().await.expect("agent shuts down");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 6)]
async fn remembered_entities_rebalance_recipe() {
use std::{
collections::BTreeMap,
net::SocketAddr,
sync::{Arc, Mutex},
};
use datum_agent::{
ClusterAgent, ClusterAgentConfig, ClusterAgentHandle, NodeSessionConfig,
dcp::{DcpJobFactories, DcpServerConfig, DcpTcpServerConfig},
};
use datum_cluster::{ClusterConfig, MemberState};
use datum_cluster_sharding::{
EntityContext, InMemoryStore, RememberEntitiesConfig, ReplyPort, Sharding, ShardingConfig,
ShardingHandle, ShardingResult,
};
#[derive(Serialize, Deserialize)]
enum CartMsg {
Touch(ReplyPort<String>),
}
fn agent_config(node_id: &str, seed_nodes: Vec<SocketAddr>) -> ClusterAgentConfig {
ClusterAgentConfig {
cluster: ClusterConfig {
node_id: node_id.to_owned(),
seed_nodes,
bind_addr: "127.0.0.1:0".parse().expect("bind addr"),
advertise_addr: "127.0.0.1:0".parse().expect("advertise addr"),
gossip_interval: Duration::from_millis(60),
probe_timeout: Duration::from_millis(15),
downing_timeout: Duration::from_millis(150),
..ClusterConfig::default()
},
dcp: DcpServerConfig {
node_id: node_id.to_owned(),
tcp: Some(DcpTcpServerConfig {
addr: "127.0.0.1:0".parse().expect("dcp addr"),
}),
..DcpServerConfig::default()
},
sessions: NodeSessionConfig {
request_timeout: Duration::from_secs(2),
command_buffer: 256,
..NodeSessionConfig::default()
},
..ClusterAgentConfig::default()
}
}
async fn wait_all_up(agents: &[&ClusterAgentHandle], expected: usize) {
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
while tokio::time::Instant::now() < deadline {
if agents.iter().all(|agent| {
agent
.cluster()
.current_state()
.members
.values()
.filter(|member| member.state == MemberState::Up && !member.unreachable)
.count()
== expected
}) {
return;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
panic!("cluster did not converge to {expected} Up members");
}
async fn register_cart(
sharding: &ShardingHandle,
starts: Arc<Mutex<BTreeMap<String, Vec<String>>>>,
) {
let node_id = sharding.node_id().to_owned();
sharding
.register_remembered_entity_type("cart", move |context: EntityContext| {
starts
.lock()
.expect("starts lock")
.entry(context.entity_id.clone())
.or_default()
.push(node_id.clone());
let reply_node = node_id.clone();
move |_context: &EntityContext, message: CartMsg| -> ShardingResult<()> {
let CartMsg::Touch(reply) = message;
let _ = reply.send(reply_node.clone());
Ok(())
}
})
.await
.expect("remembered entity type registers");
}
let starts = Arc::new(Mutex::new(BTreeMap::<String, Vec<String>>::new()));
let remember_store = Arc::new(InMemoryStore::new());
let sharding_config = ShardingConfig {
num_shards: 16,
rebalance_per_round: 16,
coordinator_tick: Duration::from_millis(25),
remember_entities: RememberEntitiesConfig::with_store(remember_store),
..ShardingConfig::default()
};
let first_agent =
ClusterAgent::start(agent_config("node-a", Vec::new()), DcpJobFactories::new())
.await
.expect("first node starts");
let first = Sharding::init(&first_agent, sharding_config.clone()).expect("first sharding");
register_cart(&first, Arc::clone(&starts)).await;
let mut entity_shards = BTreeMap::new();
for index in 0..64 {
let entity_id = format!("cart-{index}");
let entity_ref = first.entity_ref::<CartMsg>("cart", entity_id.clone());
entity_shards.insert(entity_id, entity_ref.shard_id().to_owned());
let _owner = entity_ref
.ask(Duration::from_secs(5), CartMsg::Touch)
.await
.expect("initial cart touch");
}
first
.flush_remember_entities()
.await
.expect("remembered starts flush");
let initial_table = first
.allocation_table("cart")
.await
.expect("initial allocation table");
let initial_owner_by_shard = initial_table
.entries
.iter()
.map(|entry| (entry.shard_id.clone(), entry.node_id.clone()))
.collect::<BTreeMap<_, _>>();
let second_agent = ClusterAgent::start(
agent_config("node-b", vec![first_agent.cluster().advertise_addr()]),
DcpJobFactories::new(),
)
.await
.expect("second node starts");
wait_all_up(&[&first_agent, &second_agent], 2).await;
let second = Sharding::init(&second_agent, sharding_config).expect("second sharding");
register_cart(&second, Arc::clone(&starts)).await;
let moved_entity = {
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
loop {
let table = first
.allocation_table("cart")
.await
.expect("current allocation table");
let final_owner_by_shard = table
.entries
.iter()
.map(|entry| (entry.shard_id.clone(), entry.node_id.clone()))
.collect::<BTreeMap<_, _>>();
if let Some(entity_id) = entity_shards.iter().find_map(|(entity_id, shard_id)| {
let before = initial_owner_by_shard.get(shard_id)?;
let after = final_owner_by_shard.get(shard_id)?;
(before != after).then(|| entity_id.clone())
}) {
break entity_id;
}
if tokio::time::Instant::now() >= deadline {
panic!("no remembered cart moved during rebalance");
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
};
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
while tokio::time::Instant::now() < deadline {
if starts
.lock()
.expect("starts lock")
.get(&moved_entity)
.is_some_and(|nodes| nodes.len() >= 2)
{
break;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
let started_on = starts
.lock()
.expect("starts lock")
.get(&moved_entity)
.cloned()
.expect("moved entity was started");
println!("{moved_entity} started on {started_on:?}");
assert!(started_on.len() >= 2);
first.shutdown().await;
second.shutdown().await;
first_agent.shutdown().await.expect("first shuts down");
second_agent.shutdown().await.expect("second shuts down");
}