#![forbid(unsafe_code)]
use std::{
cell::RefCell,
collections::{BTreeMap, BTreeSet, VecDeque},
fmt,
fs::{self, File, OpenOptions},
future::Future,
io::{BufRead, BufReader, Write},
marker::PhantomData,
panic::{AssertUnwindSafe, catch_unwind},
path::{Path, PathBuf},
pin::Pin,
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
use datum_agent::{
AGENT_ROLE, ClusterAgentHandle, NodeSessionManagerHandle,
dcp::{
CompleteShardingAsk, ForwardShardEnvelopes, RememberShardAllocations, ResponseStatus,
ShardAllocation, ShardAllocationEntry, ShardAllocationRequest, ShardAllocationTable,
ShardEnvelopeAck, ShardEnvelopeBatchResult, ShardEnvelopeWire, ShardingViewProvider,
},
};
use datum_cluster::{ClusterState, Member, MemberState, Signal};
use ractor::{Actor, ActorProcessingErr, ActorRef};
use serde::{Deserialize, Deserializer, Serialize, Serializer, de::DeserializeOwned};
use tokio::{
sync::{RwLock, broadcast, mpsc, oneshot},
task::JoinHandle,
time::Instant,
};
pub const VERSION: &str = env!("CARGO_PKG_VERSION");
type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
type PendingAskMap = Arc<Mutex<BTreeMap<u64, oneshot::Sender<Result<Vec<u8>, ShardingError>>>>>;
type EntityRegistry = Arc<RwLock<BTreeMap<String, Arc<dyn DynEntityType>>>>;
type AllocationCache = Arc<RwLock<BTreeMap<ShardKey, Allocation>>>;
type MovingShardSet = Arc<RwLock<BTreeSet<ShardKey>>>;
type FailedRememberShards = Arc<RwLock<BTreeMap<ShardKey, String>>>;
pub type RememberEntitiesStoreFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub type RememberEntitiesStoreResult<T> = Result<T, RememberEntitiesStoreError>;
pub type ShardingResult<T> = Result<T, ShardingError>;
#[derive(Debug, thiserror::Error)]
pub enum ShardingError {
#[error("invalid sharding config: {0}")]
InvalidConfig(String),
#[error("sharding region stopped")]
Stopped,
#[error("no eligible shard owner: {0}")]
NoEligibleShardOwner(String),
#[error("entity type is not registered: {0}")]
EntityTypeNotRegistered(String),
#[error("shard allocation buffer overflow for {type_name}/{shard_id}")]
BufferOverflow { type_name: String, shard_id: String },
#[error("sharding request timed out after {0:?}")]
Timeout(Duration),
#[error("sharding codec failed: {0}")]
Codec(String),
#[error("node session failed: {0}")]
Session(String),
#[error("entity actor failed: {0}")]
Actor(String),
#[error("remember-entities store failed for {type_name}/{shard_id}: {message}")]
RememberEntities {
type_name: String,
shard_id: String,
message: String,
},
}
impl ShardingError {
fn codec(error: impl std::fmt::Display) -> Self {
Self::Codec(error.to_string())
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("{message}")]
pub struct RememberEntitiesStoreError {
message: String,
}
impl RememberEntitiesStoreError {
#[must_use]
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
#[must_use]
pub fn message(&self) -> &str {
&self.message
}
}
impl From<std::io::Error> for RememberEntitiesStoreError {
fn from(error: std::io::Error) -> Self {
Self::new(error.to_string())
}
}
pub trait RememberEntitiesStore: Send + Sync + 'static {
fn record_entity_started<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>>;
fn record_entity_stopped<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>>;
fn list_entities_for_shard<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>>;
fn list_shards<'a>(
&'a self,
_type_name: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>> {
Box::pin(async { Ok(Vec::new()) })
}
fn flush<'a>(&'a self) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
Box::pin(async { Ok(()) })
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum RememberEntitiesFailurePolicy {
#[default]
FailOpen,
FailClosed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RememberEntitiesOperation {
RecordStarted,
RecordStopped,
ListEntities,
ListShards,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RememberEntitiesEvent {
pub type_name: String,
pub shard_id: String,
pub entity_id: Option<String>,
pub operation: RememberEntitiesOperation,
pub policy: RememberEntitiesFailurePolicy,
pub message: String,
}
#[derive(Clone)]
pub struct RememberEntitiesConfig {
store: Option<Arc<dyn RememberEntitiesStore>>,
pub queue_capacity: usize,
pub failure_policy: RememberEntitiesFailurePolicy,
pub event_buffer: usize,
}
impl fmt::Debug for RememberEntitiesConfig {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RememberEntitiesConfig")
.field("enabled", &self.store.is_some())
.field("queue_capacity", &self.queue_capacity)
.field("failure_policy", &self.failure_policy)
.field("event_buffer", &self.event_buffer)
.finish()
}
}
impl Default for RememberEntitiesConfig {
fn default() -> Self {
Self::disabled()
}
}
impl RememberEntitiesConfig {
#[must_use]
pub const fn disabled() -> Self {
Self {
store: None,
queue_capacity: 4096,
failure_policy: RememberEntitiesFailurePolicy::FailOpen,
event_buffer: 1024,
}
}
#[must_use]
pub fn with_store(store: Arc<dyn RememberEntitiesStore>) -> Self {
Self {
store: Some(store),
..Self::disabled()
}
}
#[must_use]
pub fn is_enabled(&self) -> bool {
self.store.is_some()
}
#[must_use]
pub const fn with_queue_capacity(mut self, capacity: usize) -> Self {
self.queue_capacity = capacity;
self
}
#[must_use]
pub const fn with_failure_policy(mut self, policy: RememberEntitiesFailurePolicy) -> Self {
self.failure_policy = policy;
self
}
#[must_use]
pub const fn with_event_buffer(mut self, capacity: usize) -> Self {
self.event_buffer = capacity;
self
}
}
#[derive(Debug, Default, Clone)]
pub struct InMemoryStore {
entities: Arc<Mutex<BTreeMap<RememberShardKey, BTreeSet<String>>>>,
}
impl InMemoryStore {
#[must_use]
pub fn new() -> Self {
Self::default()
}
}
impl RememberEntitiesStore for InMemoryStore {
fn record_entity_started<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
Box::pin(async move {
let mut locked = self
.entities
.lock()
.map_err(|_| RememberEntitiesStoreError::new("in-memory store poisoned"))?;
locked
.entry(RememberShardKey::new(type_name, shard_id))
.or_default()
.insert(entity_id.to_owned());
Ok(())
})
}
fn record_entity_stopped<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
Box::pin(async move {
let mut locked = self
.entities
.lock()
.map_err(|_| RememberEntitiesStoreError::new("in-memory store poisoned"))?;
let key = RememberShardKey::new(type_name, shard_id);
if let Some(entities) = locked.get_mut(&key) {
entities.remove(entity_id);
if entities.is_empty() {
locked.remove(&key);
}
}
Ok(())
})
}
fn list_entities_for_shard<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>> {
Box::pin(async move {
let locked = self
.entities
.lock()
.map_err(|_| RememberEntitiesStoreError::new("in-memory store poisoned"))?;
Ok(locked
.get(&RememberShardKey::new(type_name, shard_id))
.map(|entities| entities.iter().cloned().collect())
.unwrap_or_default())
})
}
fn list_shards<'a>(
&'a self,
type_name: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>> {
Box::pin(async move {
let locked = self
.entities
.lock()
.map_err(|_| RememberEntitiesStoreError::new("in-memory store poisoned"))?;
Ok(locked
.keys()
.filter(|key| key.type_name == type_name)
.map(|key| key.shard_id.clone())
.collect())
})
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum FileStoreFsyncPolicy {
Never,
OnCompaction,
#[default]
Always,
}
impl FileStoreFsyncPolicy {
const fn sync_append(self) -> bool {
matches!(self, Self::Always)
}
const fn sync_compaction(self) -> bool {
matches!(self, Self::OnCompaction | Self::Always)
}
}
#[derive(Debug, Clone)]
pub struct FileStore {
dir: PathBuf,
fsync_policy: FileStoreFsyncPolicy,
}
impl FileStore {
#[must_use]
pub fn new(dir: impl Into<PathBuf>) -> Self {
Self {
dir: dir.into(),
fsync_policy: FileStoreFsyncPolicy::default(),
}
}
#[must_use]
pub const fn with_fsync_policy(mut self, policy: FileStoreFsyncPolicy) -> Self {
self.fsync_policy = policy;
self
}
}
impl RememberEntitiesStore for FileStore {
fn record_entity_started<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
let path = shard_log_path(&self.dir, type_name, shard_id);
let entity_id = entity_id.to_owned();
let fsync_policy = self.fsync_policy;
Box::pin(async move {
tokio::task::spawn_blocking(move || {
append_file_store_record(&path, 'S', &entity_id, fsync_policy)
})
.await
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?
})
}
fn record_entity_stopped<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
entity_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
let path = shard_log_path(&self.dir, type_name, shard_id);
let entity_id = entity_id.to_owned();
let fsync_policy = self.fsync_policy;
Box::pin(async move {
tokio::task::spawn_blocking(move || {
let mut entities = read_file_store_entities(&path)?;
entities.remove(&entity_id);
compact_file_store_log(&path, &entities, fsync_policy)
})
.await
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?
})
}
fn list_entities_for_shard<'a>(
&'a self,
type_name: &'a str,
shard_id: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>> {
let path = shard_log_path(&self.dir, type_name, shard_id);
Box::pin(async move {
tokio::task::spawn_blocking(move || {
read_file_store_entities(&path).map(|entities| entities.into_iter().collect())
})
.await
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?
})
}
fn list_shards<'a>(
&'a self,
type_name: &'a str,
) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<Vec<String>>> {
let type_dir = self.dir.join(hex_encode(type_name.as_bytes()));
Box::pin(async move {
tokio::task::spawn_blocking(move || list_file_store_shards(&type_dir))
.await
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?
})
}
fn flush<'a>(&'a self) -> RememberEntitiesStoreFuture<'a, RememberEntitiesStoreResult<()>> {
let dir = self.dir.clone();
Box::pin(async move {
tokio::task::spawn_blocking(move || {
let mut last_error: Option<std::io::Error> = None;
fn sync_dir(path: &std::path::Path, last_error: &mut Option<std::io::Error>) {
let Ok(entries) = std::fs::read_dir(path) else {
return;
};
for entry in entries.flatten() {
let entry_path = entry.path();
if entry_path.is_dir() {
sync_dir(&entry_path, last_error);
} else if entry_path.extension().and_then(|e| e.to_str()) == Some("log")
&& let Ok(file) = std::fs::File::open(&entry_path)
&& let Err(error) = file.sync_data()
{
*last_error = Some(error);
}
}
}
sync_dir(&dir, &mut last_error);
if let Some(error) = last_error {
Err(RememberEntitiesStoreError::new(error.to_string()))
} else {
Ok(())
}
})
.await
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct RememberShardKey {
type_name: String,
shard_id: String,
}
impl RememberShardKey {
fn new(type_name: &str, shard_id: &str) -> Self {
Self {
type_name: type_name.to_owned(),
shard_id: shard_id.to_owned(),
}
}
}
fn shard_log_path(root: &Path, type_name: &str, shard_id: &str) -> PathBuf {
root.join(hex_encode(type_name.as_bytes()))
.join(format!("{}.log", hex_encode(shard_id.as_bytes())))
}
fn append_file_store_record(
path: &Path,
op: char,
entity_id: &str,
fsync_policy: FileStoreFsyncPolicy,
) -> RememberEntitiesStoreResult<()> {
if let Some(parent) = path.parent() {
fs::create_dir_all(parent)?;
}
let mut file = OpenOptions::new().create(true).append(true).open(path)?;
writeln!(file, "{op}\t{}", hex_encode(entity_id.as_bytes()))?;
if fsync_policy.sync_append() {
file.sync_all()?;
}
Ok(())
}
fn compact_file_store_log(
path: &Path,
entities: &BTreeSet<String>,
fsync_policy: FileStoreFsyncPolicy,
) -> RememberEntitiesStoreResult<()> {
if let Some(parent) = path.parent() {
fs::create_dir_all(parent)?;
}
let tmp_path = path.with_extension("log.tmp");
{
let mut file = File::create(&tmp_path)?;
for entity_id in entities {
writeln!(file, "S\t{}", hex_encode(entity_id.as_bytes()))?;
}
if fsync_policy.sync_compaction() {
file.sync_all()?;
}
}
fs::rename(&tmp_path, path)?;
Ok(())
}
fn read_file_store_entities(path: &Path) -> RememberEntitiesStoreResult<BTreeSet<String>> {
let file = match File::open(path) {
Ok(file) => file,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
return Ok(BTreeSet::new());
}
Err(error) => return Err(error.into()),
};
let mut entities = BTreeSet::new();
for line in BufReader::new(file).lines() {
let line = line?;
if line.trim().is_empty() {
continue;
}
let Some((op, encoded_entity)) = line.split_once('\t') else {
return Err(RememberEntitiesStoreError::new(
"malformed remember-entities log record",
));
};
let entity_id = String::from_utf8(hex_decode(encoded_entity)?)
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?;
match op {
"S" => {
entities.insert(entity_id);
}
"T" => {
entities.remove(&entity_id);
}
_ => {
return Err(RememberEntitiesStoreError::new(
"unknown remember-entities log operation",
));
}
}
}
Ok(entities)
}
fn list_file_store_shards(type_dir: &Path) -> RememberEntitiesStoreResult<Vec<String>> {
let entries = match fs::read_dir(type_dir) {
Ok(entries) => entries,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(Vec::new()),
Err(error) => return Err(error.into()),
};
let mut shards = Vec::new();
for entry in entries {
let entry = entry?;
let path = entry.path();
if path.extension().and_then(|extension| extension.to_str()) != Some("log") {
continue;
}
let Some(stem) = path.file_stem().and_then(|stem| stem.to_str()) else {
continue;
};
let shard_id = String::from_utf8(hex_decode(stem)?)
.map_err(|error| RememberEntitiesStoreError::new(error.to_string()))?;
shards.push(shard_id);
}
shards.sort();
shards.dedup();
Ok(shards)
}
fn hex_encode(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut encoded = String::with_capacity(bytes.len() * 2);
for byte in bytes {
encoded.push(HEX[(byte >> 4) as usize] as char);
encoded.push(HEX[(byte & 0x0f) as usize] as char);
}
encoded
}
fn hex_decode(encoded: &str) -> RememberEntitiesStoreResult<Vec<u8>> {
let bytes = encoded.as_bytes();
if !bytes.len().is_multiple_of(2) {
return Err(RememberEntitiesStoreError::new("odd-length hex string"));
}
let mut decoded = Vec::with_capacity(bytes.len() / 2);
for pair in bytes.chunks_exact(2) {
let high = hex_value(pair[0])?;
let low = hex_value(pair[1])?;
decoded.push((high << 4) | low);
}
Ok(decoded)
}
fn hex_value(byte: u8) -> RememberEntitiesStoreResult<u8> {
match byte {
b'0'..=b'9' => Ok(byte - b'0'),
b'a'..=b'f' => Ok(byte - b'a' + 10),
b'A'..=b'F' => Ok(byte - b'A' + 10),
_ => Err(RememberEntitiesStoreError::new("invalid hex digit")),
}
}
#[derive(Debug, Clone)]
pub struct ShardingConfig {
pub num_shards: u64,
pub allocation_buffer: usize,
pub request_timeout: Duration,
pub command_buffer: usize,
pub role_constraint: Option<String>,
pub agent_role: String,
pub coordinator_tick: Duration,
pub rebalance_per_round: usize,
pub passivation_idle_timeout: Option<Duration>,
pub handoff_drain_delay: Duration,
pub remember_entities: RememberEntitiesConfig,
}
impl Default for ShardingConfig {
fn default() -> Self {
Self {
num_shards: 128,
allocation_buffer: 1024,
request_timeout: Duration::from_millis(750),
command_buffer: 65_536,
role_constraint: None,
agent_role: AGENT_ROLE.to_owned(),
coordinator_tick: Duration::from_millis(50),
rebalance_per_round: 10,
passivation_idle_timeout: None,
handoff_drain_delay: Duration::from_millis(1),
remember_entities: RememberEntitiesConfig::default(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RebalanceReason {
DeadOwner,
GracefulSpread,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ShardMovement {
pub type_name: String,
pub shard_id: String,
pub from_node: String,
pub to_node: String,
pub generation: u64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RebalanceRound {
pub round: u64,
pub reason: RebalanceReason,
pub movements: Vec<ShardMovement>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ShardEnvelope<M> {
pub entity_id: String,
pub message: M,
}
impl<M> ShardEnvelope<M> {
#[must_use]
pub fn new(entity_id: impl Into<String>, message: M) -> Self {
Self {
entity_id: entity_id.into(),
message,
}
}
}
pub trait ShardExtractor<M>: Send + Sync + 'static {
fn entity_id<'a>(&self, envelope: &'a ShardEnvelope<M>) -> &'a str {
&envelope.entity_id
}
fn shard_id(&self, entity_id: &str) -> String;
}
#[derive(Debug, Clone)]
pub struct DefaultShardExtractor {
num_shards: u64,
}
impl DefaultShardExtractor {
#[must_use]
pub fn new(num_shards: u64) -> Self {
Self {
num_shards: num_shards.max(1),
}
}
#[must_use]
pub const fn num_shards(&self) -> u64 {
self.num_shards
}
}
impl<M> ShardExtractor<M> for DefaultShardExtractor {
fn shard_id(&self, entity_id: &str) -> String {
(fnv1a_64(entity_id.as_bytes()) % self.num_shards).to_string()
}
}
const FNV_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
const FNV_PRIME: u64 = 0x100000001b3;
fn fnv1a_64(data: &[u8]) -> u64 {
let mut hash = FNV_OFFSET_BASIS;
for &byte in data {
hash ^= byte as u64;
hash = hash.wrapping_mul(FNV_PRIME);
}
hash
}
#[derive(Debug, Clone)]
pub struct EntityContext {
pub type_name: String,
pub entity_id: String,
}
pub trait EntityBehavior<M>: Send + 'static {
fn handle(&mut self, context: &EntityContext, message: M) -> ShardingResult<()>;
}
impl<M, F> EntityBehavior<M> for F
where
F: FnMut(&EntityContext, M) -> ShardingResult<()> + Send + 'static,
{
fn handle(&mut self, context: &EntityContext, message: M) -> ShardingResult<()> {
self(context, message)
}
}
pub struct ReplyPort<T> {
target: WireReplyTarget,
local_node: String,
commands: mpsc::Sender<RegionCommand>,
sessions: NodeSessionManagerHandle,
pending_asks: PendingAskMap,
_marker: PhantomData<T>,
}
impl<T> std::fmt::Debug for ReplyPort<T> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ReplyPort")
.field("origin_node", &self.target.origin_node)
.field("request_id", &self.target.request_id)
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ReplySendError;
impl std::fmt::Display for ReplySendError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("sharding reply receiver dropped")
}
}
impl std::error::Error for ReplySendError {}
impl<T> ReplyPort<T> {
fn new(
target: WireReplyTarget,
local_node: String,
commands: mpsc::Sender<RegionCommand>,
sessions: NodeSessionManagerHandle,
pending_asks: PendingAskMap,
) -> Self {
Self {
target,
local_node,
commands,
sessions,
pending_asks,
_marker: PhantomData,
}
}
}
impl<T> ReplyPort<T>
where
T: Serialize + Send + 'static,
{
pub fn send(self, reply: T) -> Result<(), ReplySendError> {
let payload = encode(&reply).map_err(|_error| ReplySendError)?;
if self.target.origin_node == self.local_node {
complete_pending_ask(
&self.pending_asks,
self.target.request_id,
true,
payload,
String::new(),
);
return Ok(());
}
let response = CompleteShardingAsk {
request_id: self.target.request_id,
ok: true,
payload: payload.clone(),
message: String::new(),
};
if self
.sessions
.try_complete_sharding_ask_pipe(&self.target.origin_node, response)
.is_ok()
{
return Ok(());
}
self.commands
.try_send(RegionCommand::SendAskReply {
target: self.target,
ok: true,
payload,
message: String::new(),
})
.map_err(|_error| ReplySendError)
}
}
impl<T> Serialize for ReplyPort<T> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
self.target.serialize(serializer)
}
}
impl<'de, T> Deserialize<'de> for ReplyPort<T> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let target = WireReplyTarget::deserialize(deserializer)?;
let context = decode_context().ok_or_else(|| {
serde::de::Error::custom("ReplyPort decoded outside sharding delivery context")
})?;
Ok(Self::new(
target,
context.local_node,
context.commands,
context.sessions,
context.pending_asks,
))
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
struct WireReplyTarget {
origin_node: String,
request_id: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
enum WirePayload {
Message(Vec<u8>),
Passivate,
}
const WIRE_PAYLOAD_MAGIC: &[u8] = b"\0datum-sharding-v2\0";
fn encode_wire_payload(payload: WirePayload) -> ShardingResult<Vec<u8>> {
let mut encoded = WIRE_PAYLOAD_MAGIC.to_vec();
encoded.extend(encode(&payload)?);
Ok(encoded)
}
fn decode_wire_payload(payload: Vec<u8>) -> ShardingResult<WirePayload> {
if let Some(encoded) = payload.strip_prefix(WIRE_PAYLOAD_MAGIC) {
return decode(encoded);
}
Ok(WirePayload::Message(payload))
}
#[derive(Clone)]
struct ReplyContext {
local_node: String,
commands: mpsc::Sender<RegionCommand>,
sessions: NodeSessionManagerHandle,
pending_asks: PendingAskMap,
}
thread_local! {
static DECODE_CONTEXT: RefCell<Option<ReplyContext>> = const { RefCell::new(None) };
}
fn decode_context() -> Option<ReplyContext> {
DECODE_CONTEXT.with(|context| context.borrow().clone())
}
fn with_decode_context<T>(
reply_context: ReplyContext,
f: impl FnOnce() -> ShardingResult<T>,
) -> ShardingResult<T> {
DECODE_CONTEXT.with(|context| {
let previous = context.replace(Some(reply_context));
let result = f();
context.replace(previous);
result
})
}
pub struct Sharding;
impl Sharding {
pub fn init(
cluster_handle: &ClusterAgentHandle,
config: ShardingConfig,
) -> ShardingResult<ShardingHandle> {
validate_config(&config)?;
let (commands, receiver) = mpsc::channel(config.command_buffer.max(1));
let pending_asks = Arc::new(Mutex::new(BTreeMap::new()));
let entity_types = Arc::new(RwLock::new(BTreeMap::new()));
let allocation_cache = Arc::new(RwLock::new(BTreeMap::new()));
let moving_shards = Arc::new(RwLock::new(BTreeSet::new()));
let tasks = Arc::new(Mutex::new(Vec::new()));
let (remember_events, _) = broadcast::channel(config.remember_entities.event_buffer.max(1));
let remember = config.remember_entities.store.as_ref().map(|store| {
let (runtime, task) = RememberEntitiesRuntime::start(
Arc::clone(store),
&config.remember_entities,
remember_events.clone(),
);
tasks.lock().expect("sharding tasks poisoned").push(task);
runtime
});
let state = RegionState {
self_node: cluster_handle.cluster().node_id().to_owned(),
cluster_state: cluster_handle.cluster().state(),
sessions: cluster_handle.sessions().clone(),
config: config.clone(),
entity_types: Arc::clone(&entity_types),
allocation_cache: Arc::clone(&allocation_cache),
moving_shards: Arc::clone(&moving_shards),
remember: remember.clone(),
allocations: BTreeMap::new(),
pending_allocations: BTreeMap::new(),
allocation_inflight: BTreeMap::new(),
handoffs: BTreeMap::new(),
pending_asks: Arc::clone(&pending_asks),
active_coordinator: false,
rebuilding: false,
next_generation: 1,
next_rebalance_round: 1,
rebalance_rounds: VecDeque::new(),
commands: commands.clone(),
};
let provider = Arc::new(ShardingProvider {
commands: commands.clone(),
pending_asks: Arc::clone(&pending_asks),
timeout: config.request_timeout,
});
cluster_handle.server().set_sharding_view(provider);
let manager = tokio::spawn(run_region_manager(state, receiver));
let tick_commands = commands.clone();
let tick_interval = config.coordinator_tick;
let tick = tokio::spawn(async move {
let mut interval = tokio::time::interval(tick_interval);
loop {
interval.tick().await;
if tick_commands.send(RegionCommand::Tick).await.is_err() {
break;
}
}
});
{
let mut locked = tasks.lock().expect("sharding tasks poisoned");
locked.push(manager);
locked.push(tick);
}
Ok(ShardingHandle {
commands,
sessions: cluster_handle.sessions().clone(),
entity_types,
allocation_cache,
moving_shards,
cluster_state: cluster_handle.cluster().state(),
agent_role: config.agent_role.clone(),
role_constraint: config.role_constraint.clone(),
self_node: cluster_handle.cluster().node_id().to_owned(),
default_extractor: Arc::new(DefaultShardExtractor::new(config.num_shards)),
request_timeout: config.request_timeout,
next_request_id: Arc::new(AtomicU64::new(1)),
pending_asks,
remember,
remember_events,
tasks,
})
}
}
#[derive(Clone)]
pub struct ShardingHandle {
commands: mpsc::Sender<RegionCommand>,
sessions: NodeSessionManagerHandle,
entity_types: EntityRegistry,
allocation_cache: AllocationCache,
moving_shards: MovingShardSet,
cluster_state: Signal<ClusterState>,
agent_role: String,
role_constraint: Option<String>,
self_node: String,
default_extractor: Arc<DefaultShardExtractor>,
request_timeout: Duration,
next_request_id: Arc<AtomicU64>,
pending_asks: PendingAskMap,
remember: Option<RememberEntitiesRuntime>,
remember_events: broadcast::Sender<RememberEntitiesEvent>,
tasks: Arc<Mutex<Vec<JoinHandle<()>>>>,
}
impl ShardingHandle {
#[must_use]
pub fn node_id(&self) -> &str {
&self.self_node
}
pub async fn register_entity_type<M, F, B>(
&self,
type_name: impl Into<String>,
factory: F,
) -> ShardingResult<()>
where
M: Serialize + DeserializeOwned + Send + 'static,
F: Fn(EntityContext) -> B + Send + Sync + 'static,
B: EntityBehavior<M>,
{
let type_name = type_name.into();
if type_name.trim().is_empty() {
return Err(ShardingError::InvalidConfig(
"entity type name must not be empty".to_owned(),
));
}
let runtime = EntityTypeState::<M>::new(type_name.clone(), factory);
self.register_entity_type_runtime(type_name, runtime).await
}
pub async fn register_remembered_entity_type<M, F, B>(
&self,
type_name: impl Into<String>,
factory: F,
) -> ShardingResult<()>
where
M: Serialize + DeserializeOwned + Send + 'static,
F: Fn(EntityContext) -> B + Send + Sync + 'static,
B: EntityBehavior<M>,
{
if self.remember.is_none() {
return Err(ShardingError::InvalidConfig(
"remember-entities store must be configured before registering a remembered entity type"
.to_owned(),
));
}
let type_name = type_name.into();
if type_name.trim().is_empty() {
return Err(ShardingError::InvalidConfig(
"entity type name must not be empty".to_owned(),
));
}
let runtime = EntityTypeState::<M>::new_remembered(type_name.clone(), factory);
self.register_entity_type_runtime(type_name, runtime).await
}
async fn register_entity_type_runtime<M>(
&self,
type_name: String,
runtime: EntityTypeState<M>,
) -> ShardingResult<()>
where
M: Serialize + DeserializeOwned + Send + 'static,
{
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::RegisterType {
type_name,
runtime: Arc::new(runtime),
reply,
})
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)?
}
#[must_use]
pub fn subscribe_remember_entities_events(&self) -> broadcast::Receiver<RememberEntitiesEvent> {
self.remember_events.subscribe()
}
pub async fn flush_remember_entities(&self) -> ShardingResult<()> {
if let Some(remember) = &self.remember {
remember.flush().await
} else {
Ok(())
}
}
#[must_use]
pub fn entity_ref<M>(
&self,
type_name: impl Into<String>,
entity_id: impl Into<String>,
) -> EntityRef<M>
where
M: Serialize + Send + 'static,
{
self.entity_ref_with_extractor(type_name, entity_id, Arc::clone(&self.default_extractor))
}
#[must_use]
pub fn entity_ref_with_extractor<M, E>(
&self,
type_name: impl Into<String>,
entity_id: impl Into<String>,
extractor: Arc<E>,
) -> EntityRef<M>
where
M: Serialize + Send + 'static,
E: ShardExtractor<M>,
{
let entity_id = entity_id.into();
let shard_id = extractor.shard_id(&entity_id);
EntityRef {
type_name: type_name.into(),
entity_id,
shard_id,
commands: self.commands.clone(),
sessions: self.sessions.clone(),
entity_types: Arc::clone(&self.entity_types),
allocation_cache: Arc::clone(&self.allocation_cache),
moving_shards: Arc::clone(&self.moving_shards),
cluster_state: self.cluster_state.clone(),
agent_role: self.agent_role.clone(),
role_constraint: self.role_constraint.clone(),
self_node: self.self_node.clone(),
request_timeout: self.request_timeout,
next_request_id: Arc::clone(&self.next_request_id),
pending_asks: Arc::clone(&self.pending_asks),
remember: self.remember.clone(),
_marker: PhantomData,
}
}
pub async fn allocation_table(
&self,
type_name: impl Into<String>,
) -> ShardingResult<ShardAllocationTable> {
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::GetAllocations {
type_name: type_name.into(),
reply,
})
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)?
}
pub async fn rebalance_rounds(&self) -> ShardingResult<Vec<RebalanceRound>> {
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::GetRebalanceRounds { reply })
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)
}
pub async fn shutdown(&self) {
let _ = self.flush_remember_entities().await;
let _ = self.commands.send(RegionCommand::Shutdown).await;
if let Some(remember) = &self.remember {
remember.shutdown().await;
}
let tasks = {
let mut locked = self.tasks.lock().expect("sharding tasks poisoned");
std::mem::take(&mut *locked)
};
for task in tasks {
task.abort();
}
}
#[cfg(test)]
pub(crate) async fn test_live_entity_count(
&self,
type_name: &str,
shard_id: &str,
) -> ShardingResult<usize> {
let runtime = self
.entity_types
.read()
.await
.get(type_name)
.cloned()
.ok_or_else(|| ShardingError::EntityTypeNotRegistered(type_name.to_owned()))?;
Ok(runtime
.live_entity_count_for_shard(shard_id.to_owned())
.await)
}
}
pub struct EntityRef<M> {
type_name: String,
entity_id: String,
shard_id: String,
commands: mpsc::Sender<RegionCommand>,
sessions: NodeSessionManagerHandle,
entity_types: EntityRegistry,
allocation_cache: AllocationCache,
moving_shards: MovingShardSet,
cluster_state: Signal<ClusterState>,
agent_role: String,
role_constraint: Option<String>,
self_node: String,
request_timeout: Duration,
next_request_id: Arc<AtomicU64>,
pending_asks: PendingAskMap,
remember: Option<RememberEntitiesRuntime>,
_marker: PhantomData<M>,
}
impl<M> Clone for EntityRef<M> {
fn clone(&self) -> Self {
Self {
type_name: self.type_name.clone(),
entity_id: self.entity_id.clone(),
shard_id: self.shard_id.clone(),
commands: self.commands.clone(),
sessions: self.sessions.clone(),
entity_types: Arc::clone(&self.entity_types),
allocation_cache: Arc::clone(&self.allocation_cache),
moving_shards: Arc::clone(&self.moving_shards),
cluster_state: self.cluster_state.clone(),
agent_role: self.agent_role.clone(),
role_constraint: self.role_constraint.clone(),
self_node: self.self_node.clone(),
request_timeout: self.request_timeout,
next_request_id: Arc::clone(&self.next_request_id),
pending_asks: Arc::clone(&self.pending_asks),
remember: self.remember.clone(),
_marker: PhantomData,
}
}
}
impl<M> EntityRef<M>
where
M: Serialize + Send + 'static,
{
#[must_use]
pub fn entity_id(&self) -> &str {
&self.entity_id
}
#[must_use]
pub fn shard_id(&self) -> &str {
&self.shard_id
}
pub async fn tell(&self, message: M) -> ShardingResult<()> {
let payload = encode(&message)?;
let key = ShardKey {
type_name: self.type_name.clone(),
shard_id: self.shard_id.clone(),
};
if !self.moving_shards.read().await.contains(&key)
&& let Some(allocation) = self.allocation_cache.read().await.get(&key).cloned()
&& self.allocation_owner_is_eligible(&allocation)
{
return self
.route_allocated(&key, allocation, self.entity_id.clone(), payload)
.await;
}
self.allocation_cache.write().await.remove(&key);
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::Route {
type_name: self.type_name.clone(),
shard_id: self.shard_id.clone(),
entity_id: self.entity_id.clone(),
payload,
reply,
})
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)?
}
pub async fn passivate(&self) -> ShardingResult<()> {
let key = ShardKey {
type_name: self.type_name.clone(),
shard_id: self.shard_id.clone(),
};
if !self.moving_shards.read().await.contains(&key)
&& let Some(allocation) = self.allocation_cache.read().await.get(&key).cloned()
&& self.allocation_owner_is_eligible(&allocation)
{
return self
.passivate_allocated(&key, allocation, self.entity_id.clone())
.await;
}
self.allocation_cache.write().await.remove(&key);
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::Passivate {
type_name: self.type_name.clone(),
shard_id: self.shard_id.clone(),
entity_id: self.entity_id.clone(),
reply,
})
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)?
}
fn allocation_owner_is_eligible(&self, allocation: &Allocation) -> bool {
if allocation.node_id == self.self_node {
return true;
}
let snapshot = self.cluster_state.get();
is_eligible_node(
&snapshot,
&allocation.node_id,
&self.agent_role,
self.role_constraint.as_deref(),
)
}
async fn route_allocated(
&self,
key: &ShardKey,
allocation: Allocation,
entity_id: String,
payload: Vec<u8>,
) -> ShardingResult<()> {
if allocation.node_id == self.self_node {
if type_remembers_from_registry(&self.entity_types, &key.type_name).await
&& let Some(remember) = &self.remember
{
remember.ensure_shard_open(key).await?;
}
let record_entity_id = entity_id.clone();
let delivery = deliver_from_registry(
&self.entity_types,
&key.type_name,
key.shard_id.clone(),
entity_id,
payload,
ReplyContext {
local_node: self.self_node.clone(),
commands: self.commands.clone(),
sessions: self.sessions.clone(),
pending_asks: Arc::clone(&self.pending_asks),
},
)
.await?;
if delivery.remember_entities
&& delivery.entity_started
&& let Some(remember) = &self.remember
{
remember.record_started(key, &record_entity_id).await?;
}
return Ok(());
}
let batch = ForwardShardEnvelopes {
type_name: key.type_name.clone(),
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id.clone(),
payload: encode_wire_payload(WirePayload::Message(payload))?,
}],
};
self.sessions
.forward_shard_pipe_envelopes(&allocation.node_id, batch)
.await
.map_err(ShardingError::Session)
}
async fn passivate_allocated(
&self,
key: &ShardKey,
allocation: Allocation,
entity_id: String,
) -> ShardingResult<()> {
if allocation.node_id == self.self_node {
let (remember_entities, stored_shard_id) =
passivate_from_registry(&self.entity_types, &key.type_name, &entity_id).await?;
if remember_entities && let Some(remember) = &self.remember {
let key = ShardKey {
type_name: key.type_name.clone(),
shard_id: stored_shard_id.unwrap_or_else(|| key.shard_id.clone()),
};
remember.record_stopped(&key, &entity_id).await?;
}
return Ok(());
}
let batch = ForwardShardEnvelopes {
type_name: key.type_name.clone(),
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id.clone(),
payload: encode_wire_payload(WirePayload::Passivate)?,
}],
};
self.sessions
.forward_shard_pipe_envelopes(&allocation.node_id, batch)
.await
.map_err(ShardingError::Session)
}
pub async fn ask<R, F>(&self, timeout: Duration, make_message: F) -> ShardingResult<R>
where
R: Serialize + DeserializeOwned + Send + 'static,
F: FnOnce(ReplyPort<R>) -> M,
{
let request_id = self.next_request_id.fetch_add(1, Ordering::Relaxed);
let (sender, receiver) = oneshot::channel();
self.pending_asks
.lock()
.expect("pending asks poisoned")
.insert(request_id, sender);
let port = ReplyPort::new(
WireReplyTarget {
origin_node: self.self_node.clone(),
request_id,
},
self.self_node.clone(),
self.commands.clone(),
self.sessions.clone(),
Arc::clone(&self.pending_asks),
);
let message = make_message(port);
if let Err(error) = self.tell(message).await {
self.pending_asks
.lock()
.expect("pending asks poisoned")
.remove(&request_id);
return Err(error);
}
let timeout = if timeout.is_zero() {
self.request_timeout
} else {
timeout
};
let bytes = match tokio::time::timeout(timeout, receiver).await {
Ok(Ok(Ok(bytes))) => bytes,
Ok(Ok(Err(error))) => return Err(error),
Ok(Err(_closed)) => return Err(ShardingError::Stopped),
Err(_elapsed) => {
self.pending_asks
.lock()
.expect("pending asks poisoned")
.remove(&request_id);
return Err(ShardingError::Timeout(timeout));
}
};
decode(&bytes)
}
}
trait DynEntityType: Send + Sync {
fn remember_entities(&self) -> bool;
fn deliver(
&self,
shard_id: String,
entity_id: String,
payload: Vec<u8>,
context: ReplyContext,
) -> BoxFuture<'_, ShardingResult<DeliveryOutcome>>;
fn ensure_started(
&self,
entity_id: String,
shard_id: String,
) -> BoxFuture<'_, ShardingResult<bool>>;
fn passivate(&self, entity_id: String) -> BoxFuture<'_, ShardingResult<Option<String>>>;
fn passivate_idle(
&self,
idle_timeout: Duration,
now: Instant,
) -> BoxFuture<'_, Vec<PassivatedEntity>>;
fn stop_entities_for_shard(&self, shard_id: String) -> BoxFuture<'_, ()>;
#[cfg(test)]
fn live_entity_count_for_shard(&self, shard_id: String) -> BoxFuture<'_, usize>;
}
struct EntityTypeState<M> {
type_name: String,
remember_entities: bool,
factory: Arc<dyn Fn(EntityContext) -> Box<dyn EntityBehavior<M>> + Send + Sync>,
entities: tokio::sync::Mutex<BTreeMap<String, EntityCell<M>>>,
}
struct EntityCell<M> {
actor: ActorRef<M>,
handle: ractor::concurrency::JoinHandle<()>,
shard_id: String,
last_seen: Instant,
}
struct DeliveryOutcome {
entity_started: bool,
}
struct PassivatedEntity {
entity_id: String,
shard_id: String,
}
impl<M> EntityTypeState<M>
where
M: Serialize + DeserializeOwned + Send + 'static,
{
fn new<F, B>(type_name: String, factory: F) -> Self
where
F: Fn(EntityContext) -> B + Send + Sync + 'static,
B: EntityBehavior<M>,
{
Self::with_remember_entities(type_name, factory, false)
}
fn new_remembered<F, B>(type_name: String, factory: F) -> Self
where
F: Fn(EntityContext) -> B + Send + Sync + 'static,
B: EntityBehavior<M>,
{
Self::with_remember_entities(type_name, factory, true)
}
fn with_remember_entities<F, B>(type_name: String, factory: F, remember_entities: bool) -> Self
where
F: Fn(EntityContext) -> B + Send + Sync + 'static,
B: EntityBehavior<M>,
{
let factory = Arc::new(move |context: EntityContext| {
Box::new(factory(context)) as Box<dyn EntityBehavior<M>>
});
Self {
type_name,
remember_entities,
factory,
entities: tokio::sync::Mutex::new(BTreeMap::new()),
}
}
async fn actor_for(
&self,
entity_id: &str,
shard_id: &str,
) -> ShardingResult<(ActorRef<M>, bool)> {
let mut locked = self.entities.lock().await;
if let Some(cell) = locked.get_mut(entity_id) {
cell.last_seen = Instant::now();
cell.shard_id = shard_id.to_owned();
return Ok((cell.actor.clone(), false));
}
let context = EntityContext {
type_name: self.type_name.clone(),
entity_id: entity_id.to_owned(),
};
let actor = EntityActor {
factory: Arc::clone(&self.factory),
context,
};
let (actor_ref, handle) = Actor::spawn(None, actor, ())
.await
.map_err(|error| ShardingError::Actor(error.to_string()))?;
locked.insert(
entity_id.to_owned(),
EntityCell {
actor: actor_ref.clone(),
handle,
shard_id: shard_id.to_owned(),
last_seen: Instant::now(),
},
);
Ok((actor_ref, true))
}
async fn remove_actor(&self, entity_id: &str) -> Option<String> {
let removed = self.entities.lock().await.remove(entity_id);
if let Some(cell) = removed {
let shard_id = cell.shard_id;
cell.actor.stop(Some("passivated".to_owned()));
cell.handle.abort();
return Some(shard_id);
}
None
}
async fn remove_idle_actors(
&self,
idle_timeout: Duration,
now: Instant,
) -> Vec<PassivatedEntity> {
let idle = {
let locked = self.entities.lock().await;
locked
.iter()
.filter(|(_entity_id, cell)| now.duration_since(cell.last_seen) >= idle_timeout)
.map(|(entity_id, _cell)| entity_id.clone())
.collect::<Vec<_>>()
};
let mut removed = Vec::new();
for entity_id in idle {
if let Some(shard_id) = self.remove_actor(&entity_id).await {
removed.push(PassivatedEntity {
entity_id,
shard_id,
});
}
}
removed
}
async fn stop_entities_for_shard(&self, shard_id: &str) {
let to_stop = {
let locked = self.entities.lock().await;
locked
.iter()
.filter(|(_entity_id, cell)| cell.shard_id == shard_id)
.map(|(entity_id, _cell)| entity_id.clone())
.collect::<Vec<_>>()
};
for entity_id in to_stop {
let removed = self.entities.lock().await.remove(&entity_id);
if let Some(cell) = removed {
cell.actor.stop(Some("handoff".to_owned()));
cell.handle.abort();
}
}
}
#[cfg(test)]
async fn live_entity_count_for_shard(&self, shard_id: &str) -> usize {
self.entities
.lock()
.await
.values()
.filter(|cell| cell.shard_id == shard_id)
.count()
}
}
impl<M> DynEntityType for EntityTypeState<M>
where
M: Serialize + DeserializeOwned + Send + 'static,
{
fn remember_entities(&self) -> bool {
self.remember_entities
}
fn deliver(
&self,
shard_id: String,
entity_id: String,
payload: Vec<u8>,
context: ReplyContext,
) -> BoxFuture<'_, ShardingResult<DeliveryOutcome>> {
Box::pin(async move {
let message = with_decode_context(context, || decode::<M>(&payload))?;
let (actor, entity_started) = self.actor_for(&entity_id, &shard_id).await?;
match actor.cast(message) {
Ok(()) => Ok(DeliveryOutcome { entity_started }),
Err(_error) => {
self.remove_actor(&entity_id).await;
Err(ShardingError::Actor(format!(
"entity actor unavailable: {}/{}",
self.type_name, entity_id
)))
}
}
})
}
fn ensure_started(
&self,
entity_id: String,
shard_id: String,
) -> BoxFuture<'_, ShardingResult<bool>> {
Box::pin(async move {
let (_actor, started) = self.actor_for(&entity_id, &shard_id).await?;
Ok(started)
})
}
fn passivate(&self, entity_id: String) -> BoxFuture<'_, ShardingResult<Option<String>>> {
Box::pin(async move { Ok(self.remove_actor(&entity_id).await) })
}
fn passivate_idle(
&self,
idle_timeout: Duration,
now: Instant,
) -> BoxFuture<'_, Vec<PassivatedEntity>> {
Box::pin(async move { self.remove_idle_actors(idle_timeout, now).await })
}
fn stop_entities_for_shard(&self, shard_id: String) -> BoxFuture<'_, ()> {
Box::pin(async move { self.stop_entities_for_shard(&shard_id).await })
}
#[cfg(test)]
fn live_entity_count_for_shard(&self, shard_id: String) -> BoxFuture<'_, usize> {
Box::pin(async move { self.live_entity_count_for_shard(&shard_id).await })
}
}
struct RegistryDelivery {
remember_entities: bool,
entity_started: bool,
}
async fn deliver_from_registry(
entity_types: &EntityRegistry,
type_name: &str,
shard_id: String,
entity_id: String,
payload: Vec<u8>,
context: ReplyContext,
) -> ShardingResult<RegistryDelivery> {
let runtime = entity_types
.read()
.await
.get(type_name)
.cloned()
.ok_or_else(|| ShardingError::EntityTypeNotRegistered(type_name.to_owned()))?;
let remember_entities = runtime.remember_entities();
let outcome = runtime
.deliver(shard_id, entity_id, payload, context)
.await?;
Ok(RegistryDelivery {
remember_entities,
entity_started: outcome.entity_started,
})
}
async fn ensure_started_from_registry(
entity_types: &EntityRegistry,
type_name: &str,
shard_id: String,
entity_id: String,
) -> ShardingResult<bool> {
let runtime = entity_types
.read()
.await
.get(type_name)
.cloned()
.ok_or_else(|| ShardingError::EntityTypeNotRegistered(type_name.to_owned()))?;
if !runtime.remember_entities() {
return Ok(false);
}
runtime.ensure_started(entity_id, shard_id).await
}
async fn type_remembers_from_registry(entity_types: &EntityRegistry, type_name: &str) -> bool {
entity_types
.read()
.await
.get(type_name)
.is_some_and(|runtime| runtime.remember_entities())
}
async fn passivate_from_registry(
entity_types: &EntityRegistry,
type_name: &str,
entity_id: &str,
) -> ShardingResult<(bool, Option<String>)> {
let runtime = entity_types
.read()
.await
.get(type_name)
.cloned()
.ok_or_else(|| ShardingError::EntityTypeNotRegistered(type_name.to_owned()))?;
let remember_entities = runtime.remember_entities();
let shard_id = runtime.passivate(entity_id.to_owned()).await?;
Ok((remember_entities, shard_id))
}
struct EntityActor<M> {
factory: Arc<dyn Fn(EntityContext) -> Box<dyn EntityBehavior<M>> + Send + Sync>,
context: EntityContext,
}
struct EntityActorState<M> {
factory: Arc<dyn Fn(EntityContext) -> Box<dyn EntityBehavior<M>> + Send + Sync>,
context: EntityContext,
behavior: Box<dyn EntityBehavior<M>>,
}
impl<M> Actor for EntityActor<M>
where
M: Send + 'static,
{
type Msg = M;
type State = EntityActorState<M>;
type Arguments = ();
async fn pre_start(
&self,
_myself: ActorRef<Self::Msg>,
_args: Self::Arguments,
) -> Result<Self::State, ActorProcessingErr> {
Ok(EntityActorState {
factory: Arc::clone(&self.factory),
context: self.context.clone(),
behavior: (self.factory)(self.context.clone()),
})
}
async fn handle(
&self,
_myself: ActorRef<Self::Msg>,
message: Self::Msg,
state: &mut Self::State,
) -> Result<(), ActorProcessingErr> {
let result = catch_unwind(AssertUnwindSafe(|| {
state.behavior.handle(&state.context, message)
}));
match result {
Ok(Ok(())) => Ok(()),
Ok(Err(_)) | Err(_) => {
state.behavior = (state.factory)(state.context.clone());
Ok(())
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct ShardKey {
type_name: String,
shard_id: String,
}
#[derive(Clone)]
struct RememberEntitiesRuntime {
store: Arc<dyn RememberEntitiesStore>,
writer: mpsc::Sender<RememberWrite>,
failed_shards: FailedRememberShards,
events: broadcast::Sender<RememberEntitiesEvent>,
policy: RememberEntitiesFailurePolicy,
}
enum RememberWrite {
Started { key: ShardKey, entity_id: String },
Stopped { key: ShardKey, entity_id: String },
Flush { reply: oneshot::Sender<()> },
Shutdown { reply: oneshot::Sender<()> },
}
impl RememberEntitiesRuntime {
fn start(
store: Arc<dyn RememberEntitiesStore>,
config: &RememberEntitiesConfig,
events: broadcast::Sender<RememberEntitiesEvent>,
) -> (Self, JoinHandle<()>) {
let (writer, receiver) = mpsc::channel(config.queue_capacity.max(1));
let failed_shards = Arc::new(RwLock::new(BTreeMap::new()));
let runtime = Self {
store: Arc::clone(&store),
writer,
failed_shards: Arc::clone(&failed_shards),
events: events.clone(),
policy: config.failure_policy,
};
let task = tokio::spawn(run_remember_writer(
store,
receiver,
failed_shards,
events,
config.failure_policy,
));
(runtime, task)
}
async fn ensure_shard_open(&self, key: &ShardKey) -> ShardingResult<()> {
if self.policy == RememberEntitiesFailurePolicy::FailOpen {
return Ok(());
}
if let Some(message) = self.failed_shards.read().await.get(key).cloned() {
return Err(ShardingError::RememberEntities {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
message,
});
}
Ok(())
}
async fn record_started(&self, key: &ShardKey, entity_id: &str) -> ShardingResult<()> {
self.ensure_shard_open(key).await?;
self.try_enqueue(
key,
Some(entity_id),
RememberEntitiesOperation::RecordStarted,
RememberWrite::Started {
key: key.clone(),
entity_id: entity_id.to_owned(),
},
)
.await
}
async fn record_stopped(&self, key: &ShardKey, entity_id: &str) -> ShardingResult<()> {
self.ensure_shard_open(key).await?;
self.enqueue_blocking(
key,
Some(entity_id),
RememberEntitiesOperation::RecordStopped,
RememberWrite::Stopped {
key: key.clone(),
entity_id: entity_id.to_owned(),
},
)
.await
}
async fn list_entities(&self, key: &ShardKey) -> ShardingResult<Vec<String>> {
self.ensure_shard_open(key).await?;
match self
.store
.list_entities_for_shard(&key.type_name, &key.shard_id)
.await
{
Ok(entities) => Ok(entities),
Err(error) => {
self.report_failure(
key,
None,
RememberEntitiesOperation::ListEntities,
error.to_string(),
)
.await;
Err(ShardingError::RememberEntities {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
message: error.to_string(),
})
}
}
}
async fn list_shards(&self, type_name: &str) -> ShardingResult<Vec<String>> {
match self.store.list_shards(type_name).await {
Ok(shards) => Ok(shards),
Err(error) => {
let key = ShardKey {
type_name: type_name.to_owned(),
shard_id: String::new(),
};
self.report_failure(
&key,
None,
RememberEntitiesOperation::ListShards,
error.to_string(),
)
.await;
Err(ShardingError::RememberEntities {
type_name: type_name.to_owned(),
shard_id: String::new(),
message: error.to_string(),
})
}
}
}
async fn flush(&self) -> ShardingResult<()> {
let (reply, receiver) = oneshot::channel();
self.writer
.send(RememberWrite::Flush { reply })
.await
.map_err(|_| ShardingError::Stopped)?;
receiver.await.map_err(|_| ShardingError::Stopped)?;
Ok(())
}
async fn shutdown(&self) {
let (reply, receiver) = oneshot::channel();
if self
.writer
.send(RememberWrite::Shutdown { reply })
.await
.is_ok()
{
let _ = receiver.await;
}
}
async fn try_enqueue(
&self,
key: &ShardKey,
entity_id: Option<&str>,
operation: RememberEntitiesOperation,
write: RememberWrite,
) -> ShardingResult<()> {
match self.writer.try_send(write) {
Ok(()) => Ok(()),
Err(error) => {
let message = error.to_string();
self.report_failure(
key,
entity_id.map(ToOwned::to_owned),
operation,
message.clone(),
)
.await;
if self.policy == RememberEntitiesFailurePolicy::FailClosed {
Err(ShardingError::RememberEntities {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
message,
})
} else {
Ok(())
}
}
}
}
async fn enqueue_blocking(
&self,
key: &ShardKey,
entity_id: Option<&str>,
operation: RememberEntitiesOperation,
write: RememberWrite,
) -> ShardingResult<()> {
match self.writer.send(write).await {
Ok(()) => Ok(()),
Err(_error) => {
self.report_failure(
key,
entity_id.map(ToOwned::to_owned),
operation,
"remember-entities writer stopped".to_owned(),
)
.await;
if self.policy == RememberEntitiesFailurePolicy::FailClosed {
Err(ShardingError::RememberEntities {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
message: "remember-entities writer stopped".to_owned(),
})
} else {
Ok(())
}
}
}
}
async fn report_failure(
&self,
key: &ShardKey,
entity_id: Option<String>,
operation: RememberEntitiesOperation,
message: String,
) {
if self.policy == RememberEntitiesFailurePolicy::FailClosed {
self.failed_shards
.write()
.await
.insert(key.clone(), message.clone());
}
let _ = self.events.send(RememberEntitiesEvent {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
entity_id,
operation,
policy: self.policy,
message,
});
}
}
async fn run_remember_writer(
store: Arc<dyn RememberEntitiesStore>,
mut receiver: mpsc::Receiver<RememberWrite>,
failed_shards: FailedRememberShards,
events: broadcast::Sender<RememberEntitiesEvent>,
policy: RememberEntitiesFailurePolicy,
) {
while let Some(write) = receiver.recv().await {
match write {
RememberWrite::Started { key, entity_id } => {
if let Err(error) = store
.record_entity_started(&key.type_name, &key.shard_id, &entity_id)
.await
{
report_remember_writer_failure(
&failed_shards,
&events,
policy,
key,
Some(entity_id),
RememberEntitiesOperation::RecordStarted,
error.to_string(),
)
.await;
}
}
RememberWrite::Stopped { key, entity_id } => {
if let Err(error) = store
.record_entity_stopped(&key.type_name, &key.shard_id, &entity_id)
.await
{
report_remember_writer_failure(
&failed_shards,
&events,
policy,
key,
Some(entity_id),
RememberEntitiesOperation::RecordStopped,
error.to_string(),
)
.await;
}
}
RememberWrite::Flush { reply } => {
if let Err(error) = store.flush().await {
report_remember_writer_failure(
&failed_shards,
&events,
policy,
ShardKey {
type_name: String::new(),
shard_id: String::new(),
},
None,
RememberEntitiesOperation::ListEntities,
error.to_string(),
)
.await;
}
let _ = reply.send(());
}
RememberWrite::Shutdown { reply } => {
while let Ok(write) = receiver.try_recv() {
match write {
RememberWrite::Started { key, entity_id } => {
let _ = store
.record_entity_started(&key.type_name, &key.shard_id, &entity_id)
.await;
}
RememberWrite::Stopped { key, entity_id } => {
let _ = store
.record_entity_stopped(&key.type_name, &key.shard_id, &entity_id)
.await;
}
RememberWrite::Flush { reply } | RememberWrite::Shutdown { reply } => {
let _ = reply.send(());
}
}
}
let _ = reply.send(());
break;
}
}
}
}
async fn report_remember_writer_failure(
failed_shards: &FailedRememberShards,
events: &broadcast::Sender<RememberEntitiesEvent>,
policy: RememberEntitiesFailurePolicy,
key: ShardKey,
entity_id: Option<String>,
operation: RememberEntitiesOperation,
message: String,
) {
if policy == RememberEntitiesFailurePolicy::FailClosed {
failed_shards
.write()
.await
.insert(key.clone(), message.clone());
}
let _ = events.send(RememberEntitiesEvent {
type_name: key.type_name,
shard_id: key.shard_id,
entity_id,
operation,
policy,
message,
});
}
#[derive(Debug, Clone)]
struct Allocation {
node_id: String,
generation: u64,
}
enum BufferedEnvelope {
Message {
entity_id: String,
payload: Vec<u8>,
reply: oneshot::Sender<ShardingResult<()>>,
},
Passivate {
entity_id: String,
reply: oneshot::Sender<ShardingResult<()>>,
},
}
struct HandoffState {
allocation: Allocation,
buffer: VecDeque<BufferedEnvelope>,
}
struct PendingEnvelope {
entity_id: String,
payload: Vec<u8>,
reply: oneshot::Sender<ShardingResult<()>>,
}
struct InflightAllocation {
started_at: Instant,
}
struct RegionState {
self_node: String,
cluster_state: Signal<ClusterState>,
sessions: NodeSessionManagerHandle,
config: ShardingConfig,
entity_types: EntityRegistry,
allocation_cache: AllocationCache,
moving_shards: MovingShardSet,
remember: Option<RememberEntitiesRuntime>,
allocations: BTreeMap<ShardKey, Allocation>,
pending_allocations: BTreeMap<ShardKey, VecDeque<PendingEnvelope>>,
allocation_inflight: BTreeMap<ShardKey, InflightAllocation>,
handoffs: BTreeMap<ShardKey, HandoffState>,
pending_asks: PendingAskMap,
active_coordinator: bool,
rebuilding: bool,
next_generation: u64,
next_rebalance_round: u64,
rebalance_rounds: VecDeque<RebalanceRound>,
commands: mpsc::Sender<RegionCommand>,
}
enum RegionCommand {
RegisterType {
type_name: String,
runtime: Arc<dyn DynEntityType>,
reply: oneshot::Sender<ShardingResult<()>>,
},
Route {
type_name: String,
shard_id: String,
entity_id: String,
payload: Vec<u8>,
reply: oneshot::Sender<ShardingResult<()>>,
},
Passivate {
type_name: String,
shard_id: String,
entity_id: String,
reply: oneshot::Sender<ShardingResult<()>>,
},
ForwardedEnvelope {
type_name: String,
envelope: ShardEnvelopeWire,
reply: oneshot::Sender<ShardingResult<()>>,
},
AllocateShard {
type_name: String,
shard_id: String,
reply: oneshot::Sender<ShardingResult<ShardAllocation>>,
},
AllocationResolved {
key: ShardKey,
result: ShardingResult<ShardAllocation>,
},
RememberAllocations {
table: ShardAllocationTable,
reply: oneshot::Sender<ShardingResult<()>>,
},
GetAllocations {
type_name: String,
reply: oneshot::Sender<ShardingResult<ShardAllocationTable>>,
},
GetRebalanceRounds {
reply: oneshot::Sender<Vec<RebalanceRound>>,
},
RebuildResolved {
tables: Vec<ShardAllocationTable>,
},
HandoffReady {
key: ShardKey,
},
SendAskReply {
target: WireReplyTarget,
ok: bool,
payload: Vec<u8>,
message: String,
},
Tick,
Shutdown,
}
async fn run_region_manager(mut state: RegionState, mut receiver: mpsc::Receiver<RegionCommand>) {
while let Some(command) = receiver.recv().await {
match command {
RegionCommand::RegisterType {
type_name,
runtime,
reply,
} => {
let recover_type_name = type_name.clone();
state.entity_types.write().await.insert(type_name, runtime);
state
.recover_known_remembered_shards_for_type(recover_type_name)
.await;
let _ = reply.send(Ok(()));
}
RegionCommand::Route {
type_name,
shard_id,
entity_id,
payload,
reply,
} => {
let result = state
.route(type_name, shard_id, entity_id, payload, reply)
.await;
if let Err((error, reply)) = result {
let _ = reply.send(Err(error));
}
}
RegionCommand::Passivate {
type_name,
shard_id,
entity_id,
reply,
} => {
let result = state.passivate(type_name, shard_id, entity_id, reply).await;
if let Err((error, reply)) = result {
let _ = reply.send(Err(error));
}
}
RegionCommand::ForwardedEnvelope {
type_name,
envelope,
reply,
} => {
let result = state.forwarded_envelope(type_name, envelope, reply).await;
if let Err((error, reply)) = result {
let _ = reply.send(Err(error));
}
}
RegionCommand::AllocateShard {
type_name,
shard_id,
reply,
} => {
let result = state.allocate_for_coordinator(type_name, shard_id).await;
let _ = reply.send(result);
}
RegionCommand::AllocationResolved { key, result } => {
state.apply_allocation_result(key, result).await;
}
RegionCommand::RememberAllocations { table, reply } => {
state.merge_table(table).await;
let _ = reply.send(Ok(()));
}
RegionCommand::GetAllocations { type_name, reply } => {
let _ = reply.send(Ok(state.table_for_type(&type_name)));
}
RegionCommand::GetRebalanceRounds { reply } => {
let rounds: Vec<_> = state.rebalance_rounds.iter().cloned().collect();
let _ = reply.send(rounds);
}
RegionCommand::RebuildResolved { tables } => {
for table in tables {
state.merge_table(table).await;
}
state.rebuilding = false;
}
RegionCommand::HandoffReady { key } => {
state.drain_handoff(key).await;
}
RegionCommand::SendAskReply {
target,
ok,
payload,
message,
} => {
state.send_ask_reply(target, ok, payload, message).await;
}
RegionCommand::Tick => state.tick().await,
RegionCommand::Shutdown => break,
}
}
}
impl RegionState {
async fn route(
&mut self,
type_name: String,
shard_id: String,
entity_id: String,
payload: Vec<u8>,
reply: oneshot::Sender<ShardingResult<()>>,
) -> Result<(), (ShardingError, oneshot::Sender<ShardingResult<()>>)> {
if !self.entity_types.read().await.contains_key(&type_name) {
return Err((ShardingError::EntityTypeNotRegistered(type_name), reply));
}
let key = ShardKey {
type_name,
shard_id,
};
if self.handoffs.contains_key(&key) {
return self.buffer_handoff(
key,
BufferedEnvelope::Message {
entity_id,
payload,
reply,
},
);
}
if let Some(allocation) = self.allocations.get(&key).cloned() {
if !self.allocation_owner_is_eligible(&allocation) {
self.allocations.remove(&key);
self.allocation_cache.write().await.remove(&key);
} else {
if allocation.node_id == self.self_node {
let result = self.deliver_local(&key, entity_id, payload).await;
let _ = reply.send(result);
} else {
self.spawn_remote_forward(key, allocation.node_id, entity_id, payload, reply);
}
return Ok(());
}
}
let pending = self.pending_allocations.entry(key.clone()).or_default();
if pending.len() >= self.config.allocation_buffer {
return Err((
ShardingError::BufferOverflow {
type_name: key.type_name,
shard_id: key.shard_id,
},
reply,
));
}
pending.push_back(PendingEnvelope {
entity_id,
payload,
reply,
});
self.start_allocation_if_needed(key).await;
Ok(())
}
async fn passivate(
&mut self,
type_name: String,
shard_id: String,
entity_id: String,
reply: oneshot::Sender<ShardingResult<()>>,
) -> Result<(), (ShardingError, oneshot::Sender<ShardingResult<()>>)> {
if !self.entity_types.read().await.contains_key(&type_name) {
return Err((ShardingError::EntityTypeNotRegistered(type_name), reply));
}
let key = ShardKey {
type_name,
shard_id,
};
if self.handoffs.contains_key(&key) {
return self.buffer_handoff(key, BufferedEnvelope::Passivate { entity_id, reply });
}
let Some(allocation) = self.allocations.get(&key).cloned() else {
let result = self
.record_remembered_stopped_if_needed(&key, &entity_id)
.await;
let _ = reply.send(result);
return Ok(());
};
if !self.allocation_owner_is_eligible(&allocation) {
self.allocations.remove(&key);
self.allocation_cache.write().await.remove(&key);
let result = self
.record_remembered_stopped_if_needed(&key, &entity_id)
.await;
let _ = reply.send(result);
return Ok(());
}
if allocation.node_id == self.self_node {
let result = self.passivate_local(&key, &entity_id).await;
let _ = reply.send(result);
} else {
self.spawn_remote_passivate(key, allocation.node_id, entity_id, reply);
}
Ok(())
}
async fn forwarded_envelope(
&mut self,
type_name: String,
envelope: ShardEnvelopeWire,
reply: oneshot::Sender<ShardingResult<()>>,
) -> Result<(), (ShardingError, oneshot::Sender<ShardingResult<()>>)> {
let payload = match decode_wire_payload(envelope.payload) {
Ok(payload) => payload,
Err(error) => return Err((error, reply)),
};
match payload {
WirePayload::Message(payload) => {
self.route(
type_name,
envelope.shard_id,
envelope.entity_id,
payload,
reply,
)
.await
}
WirePayload::Passivate => {
self.passivate(type_name, envelope.shard_id, envelope.entity_id, reply)
.await
}
}
}
fn buffer_handoff(
&mut self,
key: ShardKey,
envelope: BufferedEnvelope,
) -> Result<(), (ShardingError, oneshot::Sender<ShardingResult<()>>)> {
let Some(handoff) = self.handoffs.get_mut(&key) else {
return match envelope {
BufferedEnvelope::Message { reply, .. }
| BufferedEnvelope::Passivate { reply, .. } => Err((ShardingError::Stopped, reply)),
};
};
if handoff.buffer.len() >= self.config.allocation_buffer {
let error = ShardingError::BufferOverflow {
type_name: key.type_name,
shard_id: key.shard_id,
};
return match envelope {
BufferedEnvelope::Message { reply, .. }
| BufferedEnvelope::Passivate { reply, .. } => Err((error, reply)),
};
}
handoff.buffer.push_back(envelope);
Ok(())
}
async fn drain_handoff(&mut self, key: ShardKey) {
let Some(mut handoff) = self.handoffs.remove(&key) else {
return;
};
let losing_owner = handoff.allocation.node_id != self.self_node;
while let Some(envelope) = handoff.buffer.pop_front() {
let _ = self
.route_buffered_to_owner(&key, &handoff.allocation.node_id, envelope)
.await;
}
if losing_owner {
let runtimes: Vec<_> = self.entity_types.read().await.values().cloned().collect();
for runtime in runtimes {
runtime.stop_entities_for_shard(key.shard_id.clone()).await;
}
}
self.moving_shards.write().await.remove(&key);
}
async fn route_buffered_to_owner(
&self,
key: &ShardKey,
owner: &str,
envelope: BufferedEnvelope,
) -> ShardingResult<()> {
match envelope {
BufferedEnvelope::Message {
entity_id,
payload,
reply,
} => {
let result = self.route_to_owner(key, owner, entity_id, payload).await;
let ok = result.is_ok();
let _ = reply.send(result);
if ok {
Ok(())
} else {
Err(ShardingError::Stopped)
}
}
BufferedEnvelope::Passivate { entity_id, reply } => {
let result = self.passivate_to_owner(key, owner, entity_id).await;
let ok = result.is_ok();
let _ = reply.send(result);
if ok {
Ok(())
} else {
Err(ShardingError::Stopped)
}
}
}
}
async fn start_handoff(&mut self, key: ShardKey, allocation: Allocation) {
self.handoffs
.entry(key.clone())
.and_modify(|handoff| {
handoff.allocation = allocation.clone();
})
.or_insert_with(|| HandoffState {
allocation,
buffer: VecDeque::new(),
});
self.moving_shards.write().await.insert(key.clone());
let commands = self.commands.clone();
let delay = self.config.handoff_drain_delay;
tokio::spawn(async move {
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
let _ = commands.send(RegionCommand::HandoffReady { key }).await;
});
}
fn allocation_owner_is_eligible(&self, allocation: &Allocation) -> bool {
if allocation.node_id == self.self_node {
return true;
}
let snapshot = self.cluster_state.get();
is_eligible_node(
&snapshot,
&allocation.node_id,
&self.config.agent_role,
self.config.role_constraint.as_deref(),
)
}
async fn drain_pending_for_allocation(&mut self, key: &ShardKey) {
let Some(allocation) = self.allocations.get(key).cloned() else {
return;
};
let pending = self.pending_allocations.remove(key).unwrap_or_default();
for envelope in pending {
let result = self
.route_to_owner(
key,
&allocation.node_id,
envelope.entity_id,
envelope.payload,
)
.await;
let _ = envelope.reply.send(result);
}
}
async fn invalidate_noneligible_allocations(&mut self) {
let stale = self
.allocations
.iter()
.filter(|(_key, allocation)| !self.allocation_owner_is_eligible(allocation))
.map(|(key, _allocation)| key.clone())
.collect::<Vec<_>>();
for key in stale {
self.allocations.remove(&key);
self.allocation_cache.write().await.remove(&key);
}
}
async fn passivate_idle_entities(&mut self) {
let Some(idle_timeout) = self.config.passivation_idle_timeout else {
return;
};
let runtimes = self
.entity_types
.read()
.await
.iter()
.map(|(type_name, runtime)| (type_name.clone(), Arc::clone(runtime)))
.collect::<Vec<_>>();
let now = Instant::now();
for (type_name, runtime) in runtimes {
let passivated = runtime.passivate_idle(idle_timeout, now).await;
if runtime.remember_entities()
&& let Some(remember) = &self.remember
{
for entity in passivated {
let key = ShardKey {
type_name: type_name.clone(),
shard_id: entity.shard_id,
};
let _ = remember.record_stopped(&key, &entity.entity_id).await;
}
}
}
}
async fn type_remembers(&self, type_name: &str) -> bool {
type_remembers_from_registry(&self.entity_types, type_name).await
}
async fn record_remembered_stopped_if_needed(
&self,
key: &ShardKey,
entity_id: &str,
) -> ShardingResult<()> {
if self.type_remembers(&key.type_name).await
&& let Some(remember) = &self.remember
{
remember.record_stopped(key, entity_id).await?;
}
Ok(())
}
async fn recover_known_remembered_shards_for_type(&mut self, type_name: String) {
if !self.type_remembers(&type_name).await {
return;
}
let Some(remember) = self.remember.clone() else {
return;
};
let Ok(shards) = remember.list_shards(&type_name).await else {
return;
};
for shard_id in shards {
self.start_allocation_if_needed(ShardKey {
type_name: type_name.clone(),
shard_id,
})
.await;
}
self.recover_local_shards_for_type(&type_name).await;
}
async fn recover_local_shards_for_type(&self, type_name: &str) {
let keys = self
.allocations
.iter()
.filter(|(key, allocation)| {
key.type_name == type_name && allocation.node_id == self.self_node
})
.map(|(key, _allocation)| key.clone())
.collect::<Vec<_>>();
for key in keys {
self.spawn_remembered_recovery(key).await;
}
}
async fn maybe_recover_local_shard(&self, key: &ShardKey, allocation: &Allocation) {
if allocation.node_id == self.self_node {
self.spawn_remembered_recovery(key.clone()).await;
}
}
async fn spawn_remembered_recovery(&self, key: ShardKey) {
if !self.type_remembers(&key.type_name).await {
return;
}
let Some(remember) = self.remember.clone() else {
return;
};
let entity_types = Arc::clone(&self.entity_types);
tokio::spawn(async move {
let Ok(entity_ids) = remember.list_entities(&key).await else {
return;
};
for entity_id in entity_ids {
let _ = ensure_started_from_registry(
&entity_types,
&key.type_name,
key.shard_id.clone(),
entity_id,
)
.await;
}
});
}
async fn start_allocation_if_needed(&mut self, key: ShardKey) {
if self.allocation_inflight.contains_key(&key) {
return;
}
self.allocation_inflight.insert(
key.clone(),
InflightAllocation {
started_at: Instant::now(),
},
);
if self.is_local_coordinator() {
let result = self
.allocate_for_coordinator(key.type_name.clone(), key.shard_id.clone())
.await;
self.apply_allocation_result(key, result).await;
return;
}
let Some(coordinator) = self.coordinator_node_id() else {
self.apply_allocation_result(
key,
Err(ShardingError::NoEligibleShardOwner(
"no coordinator is available".to_owned(),
)),
)
.await;
return;
};
let sessions = self.sessions.clone();
let commands = self.commands.clone();
let timeout = self.config.request_timeout;
tokio::spawn(async move {
let result = sessions
.allocate_shard(
&coordinator,
key.type_name.clone(),
key.shard_id.clone(),
timeout,
)
.await
.map_err(ShardingError::Session);
let _ = commands
.send(RegionCommand::AllocationResolved { key, result })
.await;
});
}
async fn apply_allocation_result(
&mut self,
key: ShardKey,
result: ShardingResult<ShardAllocation>,
) {
self.allocation_inflight.remove(&key);
match result {
Ok(allocation) => {
let allocation = Allocation {
node_id: allocation.node_id,
generation: allocation.generation,
};
self.allocations.insert(key.clone(), allocation.clone());
self.allocation_cache
.write()
.await
.insert(key.clone(), allocation.clone());
self.maybe_recover_local_shard(&key, &allocation).await;
self.drain_pending_for_allocation(&key).await;
}
Err(error) => {
let pending = self.pending_allocations.remove(&key).unwrap_or_default();
let message = error.to_string();
for envelope in pending {
let _ = envelope
.reply
.send(Err(ShardingError::Session(message.clone())));
}
}
}
}
async fn route_to_owner(
&self,
key: &ShardKey,
owner: &str,
entity_id: String,
payload: Vec<u8>,
) -> ShardingResult<()> {
if owner == self.self_node {
return self.deliver_local(key, entity_id, payload).await;
}
let batch = ForwardShardEnvelopes {
type_name: key.type_name.clone(),
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id.clone(),
payload: encode_wire_payload(WirePayload::Message(payload))?,
}],
};
self.sessions
.forward_shard_pipe_envelopes(owner, batch)
.await
.map_err(ShardingError::Session)?;
Ok(())
}
fn spawn_remote_forward(
&self,
key: ShardKey,
owner: String,
entity_id: String,
payload: Vec<u8>,
reply: oneshot::Sender<ShardingResult<()>>,
) {
let sessions = self.sessions.clone();
tokio::spawn(async move {
let batch = ForwardShardEnvelopes {
type_name: key.type_name,
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id,
payload: match encode_wire_payload(WirePayload::Message(payload)) {
Ok(payload) => payload,
Err(error) => {
let _ = reply.send(Err(error));
return;
}
},
}],
};
let result = sessions
.forward_shard_pipe_envelopes(&owner, batch)
.await
.map_err(ShardingError::Session);
let _ = reply.send(result);
});
}
async fn passivate_to_owner(
&self,
key: &ShardKey,
owner: &str,
entity_id: String,
) -> ShardingResult<()> {
if owner == self.self_node {
return self.passivate_local(key, &entity_id).await;
}
let batch = ForwardShardEnvelopes {
type_name: key.type_name.clone(),
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id.clone(),
payload: encode_wire_payload(WirePayload::Passivate)?,
}],
};
self.sessions
.forward_shard_pipe_envelopes(owner, batch)
.await
.map_err(ShardingError::Session)?;
Ok(())
}
fn spawn_remote_passivate(
&self,
key: ShardKey,
owner: String,
entity_id: String,
reply: oneshot::Sender<ShardingResult<()>>,
) {
let sessions = self.sessions.clone();
tokio::spawn(async move {
let payload = match encode_wire_payload(WirePayload::Passivate) {
Ok(payload) => payload,
Err(error) => {
let _ = reply.send(Err(error));
return;
}
};
let batch = ForwardShardEnvelopes {
type_name: key.type_name,
envelopes: vec![ShardEnvelopeWire {
entity_id,
shard_id: key.shard_id,
payload,
}],
};
let result = sessions
.forward_shard_pipe_envelopes(&owner, batch)
.await
.map_err(ShardingError::Session);
let _ = reply.send(result);
});
}
async fn deliver_local(
&self,
key: &ShardKey,
entity_id: String,
payload: Vec<u8>,
) -> ShardingResult<()> {
if self.type_remembers(&key.type_name).await
&& let Some(remember) = &self.remember
{
remember.ensure_shard_open(key).await?;
}
let record_entity_id = entity_id.clone();
let delivery = deliver_from_registry(
&self.entity_types,
&key.type_name,
key.shard_id.clone(),
entity_id,
payload,
ReplyContext {
local_node: self.self_node.clone(),
commands: self.commands.clone(),
sessions: self.sessions.clone(),
pending_asks: Arc::clone(&self.pending_asks),
},
)
.await?;
if delivery.remember_entities
&& delivery.entity_started
&& let Some(remember) = &self.remember
{
remember.record_started(key, &record_entity_id).await?;
}
Ok(())
}
async fn passivate_local(&self, key: &ShardKey, entity_id: &str) -> ShardingResult<()> {
let (remember_entities, stored_shard_id) =
passivate_from_registry(&self.entity_types, &key.type_name, entity_id).await?;
if remember_entities && let Some(remember) = &self.remember {
let key = ShardKey {
type_name: key.type_name.clone(),
shard_id: stored_shard_id.unwrap_or_else(|| key.shard_id.clone()),
};
remember.record_stopped(&key, entity_id).await?;
}
Ok(())
}
async fn allocate_for_coordinator(
&mut self,
type_name: String,
shard_id: String,
) -> ShardingResult<ShardAllocation> {
if !self.is_local_coordinator() {
return Err(ShardingError::Session(
"this node is not the sharding coordinator".to_owned(),
));
}
let key = ShardKey {
type_name: type_name.clone(),
shard_id: shard_id.clone(),
};
if let Some(allocation) = self.allocations.get(&key)
&& self.allocation_owner_is_eligible(allocation)
{
return Ok(ShardAllocation {
type_name,
shard_id,
node_id: allocation.node_id.clone(),
generation: allocation.generation,
});
}
self.allocations.remove(&key);
self.allocation_cache.write().await.remove(&key);
let target = self.choose_least_shards_owner(&type_name)?;
let generation = self.next_generation;
self.next_generation = self.next_generation.saturating_add(1).max(1);
let allocation = Allocation {
node_id: target.clone(),
generation,
};
self.allocations.insert(key.clone(), allocation.clone());
self.allocation_cache
.write()
.await
.insert(key.clone(), allocation.clone());
self.maybe_recover_local_shard(&key, &allocation).await;
self.replicate_allocations(type_name.clone()).await;
Ok(ShardAllocation {
type_name,
shard_id,
node_id: target,
generation,
})
}
fn choose_least_shards_owner(&self, type_name: &str) -> ShardingResult<String> {
let snapshot = self.cluster_state.get();
let mut candidates = eligible_members(&snapshot, &self.config)
.into_iter()
.map(|member| member.node_id.clone())
.collect::<Vec<_>>();
candidates.sort();
candidates
.into_iter()
.min_by(|left, right| {
let left_count = self.shards_on_node(type_name, left);
let right_count = self.shards_on_node(type_name, right);
left_count.cmp(&right_count).then_with(|| left.cmp(right))
})
.ok_or_else(|| {
ShardingError::NoEligibleShardOwner(format!(
"no Up members with role {}",
self.config.agent_role
))
})
}
fn shards_on_node(&self, type_name: &str, node_id: &str) -> usize {
self.allocations
.iter()
.filter(|(key, allocation)| key.type_name == type_name && allocation.node_id == node_id)
.count()
}
async fn replicate_allocations(&self, type_name: String) {
let table = self.table_for_type(&type_name);
let peers = self.peer_node_ids();
for node_id in peers {
let sessions = self.sessions.clone();
let table = table.clone();
let timeout = self.config.request_timeout.min(Duration::from_millis(250));
tokio::spawn(async move {
let _ = sessions
.remember_shard_allocations(&node_id, table, timeout)
.await;
});
}
}
fn table_for_type(&self, type_name: &str) -> ShardAllocationTable {
ShardAllocationTable {
entries: self
.allocations
.iter()
.filter(|(key, _allocation)| type_name.is_empty() || key.type_name == type_name)
.map(|(key, allocation)| ShardAllocationEntry {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
node_id: allocation.node_id.clone(),
generation: allocation.generation,
})
.collect(),
}
}
async fn merge_table(&mut self, table: ShardAllocationTable) {
for entry in table.entries {
let key = ShardKey {
type_name: entry.type_name,
shard_id: entry.shard_id,
};
let existing = self.allocations.get(&key).cloned();
let replace = existing
.as_ref()
.map(|existing| {
entry.generation > existing.generation
|| (entry.generation == existing.generation
&& entry.node_id > existing.node_id)
})
.unwrap_or(true);
if replace {
self.next_generation = self.next_generation.max(entry.generation.saturating_add(1));
let allocation = Allocation {
node_id: entry.node_id,
generation: entry.generation,
};
let owner_changed = existing
.as_ref()
.is_some_and(|existing| existing.node_id != allocation.node_id);
self.allocations.insert(key.clone(), allocation.clone());
self.allocation_cache
.write()
.await
.insert(key.clone(), allocation.clone());
self.maybe_recover_local_shard(&key, &allocation).await;
if owner_changed {
self.start_handoff(key.clone(), allocation).await;
} else {
self.drain_pending_for_allocation(&key).await;
}
}
}
}
async fn send_ask_reply(
&mut self,
target: WireReplyTarget,
ok: bool,
payload: Vec<u8>,
message: String,
) {
if target.origin_node == self.self_node {
complete_pending_ask(&self.pending_asks, target.request_id, ok, payload, message);
return;
}
let response = CompleteShardingAsk {
request_id: target.request_id,
ok,
payload,
message,
};
let _ = self
.sessions
.complete_sharding_ask_pipe(&target.origin_node, response)
.await;
}
async fn tick(&mut self) {
self.passivate_idle_entities().await;
let is_coordinator = self.is_local_coordinator();
if !is_coordinator {
self.active_coordinator = false;
self.rebuilding = false;
self.invalidate_noneligible_allocations().await;
return;
}
self.fail_stale_allocations().await;
if self.rebuilding {
return;
}
if self.active_coordinator {
self.rebalance_dead_owners().await;
self.rebalance_graceful_spread().await;
return;
}
self.active_coordinator = true;
self.rebuilding = true;
let peers = self.peer_node_ids();
let sessions = self.sessions.clone();
let commands = self.commands.clone();
let timeout = self.config.request_timeout;
let local_table = self.table_for_type("");
tokio::spawn(async move {
let mut tables = vec![local_table];
let mut pending = Vec::new();
for node_id in peers {
let sessions = sessions.clone();
pending.push(tokio::spawn(async move {
sessions
.get_shard_allocations(&node_id, String::new(), timeout)
.await
}));
}
for task in pending {
if let Ok(Ok(table)) = task.await {
tables.push(table);
}
}
let _ = commands
.send(RegionCommand::RebuildResolved { tables })
.await;
});
}
async fn rebalance_dead_owners(&mut self) {
let keys = self
.allocations
.iter()
.filter(|(_key, allocation)| !self.allocation_owner_is_eligible(allocation))
.map(|(key, _allocation)| key.clone())
.collect::<Vec<_>>();
let mut movements = Vec::new();
for key in keys {
let Ok(target) = self.choose_least_shards_owner(&key.type_name) else {
continue;
};
if let Some(movement) = self.move_shard(&key, target).await {
movements.push(movement);
}
}
self.record_rebalance_round(RebalanceReason::DeadOwner, movements)
.await;
}
async fn rebalance_graceful_spread(&mut self) {
let cap = self.config.rebalance_per_round;
if cap == 0 {
return;
}
let candidates = self.eligible_node_ids();
if candidates.len() < 2 {
return;
}
let types = self
.allocations
.keys()
.map(|key| key.type_name.clone())
.collect::<BTreeSet<_>>();
let mut movements = Vec::new();
let mut moved_keys = BTreeSet::new();
for type_name in types {
loop {
if movements.len() >= cap {
self.record_rebalance_round(RebalanceReason::GracefulSpread, movements)
.await;
return;
}
let counts = self.shard_counts_for_type(&type_name, &candidates);
let Some((source, source_count)) = counts
.iter()
.max_by(|left, right| left.1.cmp(right.1).then_with(|| left.0.cmp(right.0)))
.map(|(node, count)| (node.clone(), *count))
else {
break;
};
let Some((target, target_count)) = counts
.iter()
.min_by(|left, right| left.1.cmp(right.1).then_with(|| left.0.cmp(right.0)))
.map(|(node, count)| (node.clone(), *count))
else {
break;
};
if source_count <= target_count.saturating_add(1) {
break;
}
let Some(key) = self
.allocations
.iter()
.filter(|(key, allocation)| {
key.type_name == type_name
&& allocation.node_id == source
&& !moved_keys.contains(*key)
})
.map(|(key, _allocation)| key.clone())
.next()
else {
break;
};
moved_keys.insert(key.clone());
if let Some(movement) = self.move_shard(&key, target).await {
movements.push(movement);
} else {
break;
}
}
}
self.record_rebalance_round(RebalanceReason::GracefulSpread, movements)
.await;
}
async fn move_shard(&mut self, key: &ShardKey, target: String) -> Option<ShardMovement> {
let existing = self.allocations.get(key)?.clone();
if existing.node_id == target {
return None;
}
let generation = self.next_generation;
self.next_generation = self.next_generation.saturating_add(1).max(1);
let allocation = Allocation {
node_id: target.clone(),
generation,
};
self.allocations.insert(key.clone(), allocation.clone());
self.allocation_cache
.write()
.await
.insert(key.clone(), allocation.clone());
self.maybe_recover_local_shard(key, &allocation).await;
self.start_handoff(key.clone(), allocation).await;
Some(ShardMovement {
type_name: key.type_name.clone(),
shard_id: key.shard_id.clone(),
from_node: existing.node_id,
to_node: target,
generation,
})
}
}
fn push_bounded_round(queue: &mut VecDeque<RebalanceRound>, round: RebalanceRound) {
const MAX_ROUNDS: usize = 1024;
if queue.len() >= MAX_ROUNDS {
queue.pop_front();
}
queue.push_back(round);
}
impl RegionState {
async fn record_rebalance_round(
&mut self,
reason: RebalanceReason,
movements: Vec<ShardMovement>,
) {
if movements.is_empty() {
return;
}
let affected_types = movements
.iter()
.map(|movement| movement.type_name.clone())
.collect::<BTreeSet<_>>();
let round = RebalanceRound {
round: self.next_rebalance_round,
reason,
movements,
};
self.next_rebalance_round = self.next_rebalance_round.saturating_add(1).max(1);
push_bounded_round(&mut self.rebalance_rounds, round);
for type_name in affected_types {
self.replicate_allocations(type_name).await;
}
}
fn eligible_node_ids(&self) -> Vec<String> {
let snapshot = self.cluster_state.get();
let mut nodes = eligible_members(&snapshot, &self.config)
.into_iter()
.map(|member| member.node_id.clone())
.collect::<Vec<_>>();
nodes.sort();
nodes
}
fn shard_counts_for_type(
&self,
type_name: &str,
candidates: &[String],
) -> BTreeMap<String, usize> {
let mut counts = candidates
.iter()
.map(|node_id| (node_id.clone(), 0_usize))
.collect::<BTreeMap<_, _>>();
for (key, allocation) in &self.allocations {
if key.type_name == type_name
&& let Some(count) = counts.get_mut(&allocation.node_id)
{
*count += 1;
}
}
counts
}
async fn fail_stale_allocations(&mut self) {
let timeout = self.config.request_timeout;
let now = Instant::now();
let stale = self
.allocation_inflight
.iter()
.filter(|(_key, inflight)| now.duration_since(inflight.started_at) > timeout)
.map(|(key, _inflight)| key.clone())
.collect::<Vec<_>>();
for key in stale {
self.apply_allocation_result(key, Err(ShardingError::Timeout(timeout)))
.await;
}
}
fn peer_node_ids(&self) -> Vec<String> {
let snapshot = self.cluster_state.get();
eligible_members(&snapshot, &self.config)
.into_iter()
.filter(|member| member.node_id != self.self_node)
.map(|member| member.node_id.clone())
.collect()
}
fn is_local_coordinator(&self) -> bool {
self.cluster_state
.get()
.placement_coordinator()
.is_some_and(|member| member.node_id == self.self_node)
}
fn coordinator_node_id(&self) -> Option<String> {
self.cluster_state
.get()
.placement_coordinator()
.map(|member| member.node_id.clone())
}
}
fn eligible_members<'a>(snapshot: &'a ClusterState, config: &ShardingConfig) -> Vec<&'a Member> {
snapshot
.members
.values()
.filter(|member| member.state == MemberState::Up && !member.unreachable)
.filter(|member| member.has_role(&config.agent_role))
.filter(|member| {
config
.role_constraint
.as_deref()
.is_none_or(|role| member.has_role(role))
})
.collect()
}
fn is_eligible_node(
snapshot: &ClusterState,
node_id: &str,
agent_role: &str,
role_constraint: Option<&str>,
) -> bool {
snapshot.member(node_id).is_some_and(|member| {
member.state == MemberState::Up
&& !member.unreachable
&& member.has_role(agent_role)
&& role_constraint.is_none_or(|role| member.has_role(role))
})
}
struct ShardingProvider {
commands: mpsc::Sender<RegionCommand>,
pending_asks: PendingAskMap,
timeout: Duration,
}
impl ShardingViewProvider for ShardingProvider {
fn allocate_shard(
&self,
request: ShardAllocationRequest,
timeout: Duration,
) -> datum_agent::dcp::server::ClusterViewFuture<'_, ShardAllocation> {
Box::pin(async move {
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::AllocateShard {
type_name: request.type_name,
shard_id: request.shard_id,
reply,
})
.await
.map_err(|_| dcp_error(ResponseStatus::Failed, "sharding region stopped"))?;
wait_dcp(timeout, receiver).await
})
}
fn remember_shard_allocations(
&self,
request: RememberShardAllocations,
) -> datum_agent::dcp::server::ClusterViewFuture<'_, ()> {
Box::pin(async move {
let table = request.table.unwrap_or_else(|| ShardAllocationTable {
entries: Vec::new(),
});
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::RememberAllocations { table, reply })
.await
.map_err(|_| dcp_error(ResponseStatus::Failed, "sharding region stopped"))?;
wait_dcp(self.timeout, receiver).await
})
}
fn get_shard_allocations(
&self,
type_name: String,
timeout: Duration,
) -> datum_agent::dcp::server::ClusterViewFuture<'_, ShardAllocationTable> {
Box::pin(async move {
let (reply, receiver) = oneshot::channel();
self.commands
.send(RegionCommand::GetAllocations { type_name, reply })
.await
.map_err(|_| dcp_error(ResponseStatus::Failed, "sharding region stopped"))?;
wait_dcp(timeout, receiver).await
})
}
fn forward_shard_envelopes(
&self,
request: ForwardShardEnvelopes,
timeout: Duration,
) -> datum_agent::dcp::server::ClusterViewFuture<'_, ShardEnvelopeBatchResult> {
Box::pin(async move {
let timeout = if timeout.is_zero() {
self.timeout
} else {
timeout
};
let mut acks = Vec::with_capacity(request.envelopes.len());
for envelope in request.envelopes {
let entity_id = envelope.entity_id.clone();
let shard_id = envelope.shard_id.clone();
let (reply, receiver) = oneshot::channel();
let send_result = self
.commands
.send(RegionCommand::ForwardedEnvelope {
type_name: request.type_name.clone(),
envelope,
reply,
})
.await
.map_err(|_| dcp_error(ResponseStatus::Failed, "sharding region stopped"));
let result = match send_result {
Ok(()) => wait_dcp(timeout, receiver)
.await
.map_err(|error| error.to_string()),
Err(error) => Err(error.to_string()),
};
acks.push(ShardEnvelopeAck {
entity_id,
shard_id,
ok: result.is_ok(),
message: result.err().unwrap_or_default(),
});
}
Ok(ShardEnvelopeBatchResult { acks })
})
}
fn complete_sharding_ask(
&self,
request: CompleteShardingAsk,
) -> datum_agent::dcp::server::ClusterViewFuture<'_, ()> {
Box::pin(async move {
complete_pending_ask(
&self.pending_asks,
request.request_id,
request.ok,
request.payload,
request.message,
);
Ok(())
})
}
}
async fn wait_dcp<T>(
timeout: Duration,
receiver: oneshot::Receiver<ShardingResult<T>>,
) -> datum_agent::dcp::DcpResult<T> {
tokio::time::timeout(timeout, receiver)
.await
.map_err(|_| {
dcp_error(
ResponseStatus::DeadlineExceeded,
"sharding request timed out",
)
})?
.map_err(|_| dcp_error(ResponseStatus::Failed, "sharding region stopped"))?
.map_err(|error| dcp_error(ResponseStatus::Failed, error.to_string()))
}
fn complete_pending_ask(
pending_asks: &PendingAskMap,
request_id: u64,
ok: bool,
payload: Vec<u8>,
message: String,
) {
let sender = pending_asks
.lock()
.expect("pending asks poisoned")
.remove(&request_id);
if let Some(sender) = sender {
let result = if ok {
Ok(payload)
} else {
Err(ShardingError::Session(message))
};
let _ = sender.send(result);
}
}
fn dcp_error(status: ResponseStatus, message: impl Into<String>) -> datum_agent::dcp::DcpError {
datum_agent::dcp::DcpError::Response {
status,
message: message.into(),
}
}
fn validate_config(config: &ShardingConfig) -> ShardingResult<()> {
if config.num_shards == 0 {
return Err(ShardingError::InvalidConfig(
"num_shards must be non-zero".to_owned(),
));
}
if config.allocation_buffer == 0 {
return Err(ShardingError::InvalidConfig(
"allocation_buffer must be non-zero".to_owned(),
));
}
if config.request_timeout.is_zero() {
return Err(ShardingError::InvalidConfig(
"request_timeout must be non-zero".to_owned(),
));
}
if config.agent_role.trim().is_empty() {
return Err(ShardingError::InvalidConfig(
"agent_role must not be empty".to_owned(),
));
}
if config
.passivation_idle_timeout
.is_some_and(|timeout| timeout.is_zero())
{
return Err(ShardingError::InvalidConfig(
"passivation_idle_timeout must be non-zero when configured".to_owned(),
));
}
if config.remember_entities.is_enabled() && config.remember_entities.queue_capacity == 0 {
return Err(ShardingError::InvalidConfig(
"remember_entities.queue_capacity must be non-zero when configured".to_owned(),
));
}
if config.remember_entities.event_buffer == 0 {
return Err(ShardingError::InvalidConfig(
"remember_entities.event_buffer must be non-zero".to_owned(),
));
}
Ok(())
}
fn encode<T: Serialize>(value: &T) -> ShardingResult<Vec<u8>> {
bincode::serde::encode_to_vec(value, bincode::config::standard()).map_err(ShardingError::codec)
}
fn decode<T: DeserializeOwned>(payload: &[u8]) -> ShardingResult<T> {
let (value, read): (T, usize) =
bincode::serde::decode_from_slice(payload, bincode::config::standard())
.map_err(ShardingError::codec)?;
if read != payload.len() {
return Err(ShardingError::Codec(
"trailing bytes in sharding payload".to_owned(),
));
}
Ok(value)
}
#[cfg(test)]
mod tests {
use super::*;
fn generation_wins(
entry_gen: u64,
entry_node: &str,
existing_gen: u64,
existing_node: &str,
) -> bool {
entry_gen > existing_gen || (entry_gen == existing_gen && entry_node > existing_node)
}
#[test]
fn merge_table_generation_tie_break_converges() {
assert!(generation_wins(5, "node-b", 5, "node-a"));
assert!(!generation_wins(5, "node-a", 5, "node-b"));
assert!(generation_wins(10, "node-a", 5, "node-b"));
assert!(!generation_wins(5, "node-a", 10, "node-b"));
assert!(!generation_wins(5, "node-a", 5, "node-a"));
let (generation, a, b) = (42, "coord-0", "coord-1");
let ab = generation_wins(generation, b, generation, a);
let ba = !generation_wins(generation, a, generation, b);
assert!(ab, "b should win over a under generation(42) tie");
assert!(ba, "a should NOT win over b under same-generation tie");
assert_eq!(ab, ba);
}
#[test]
fn golden_shard_map_is_stable() {
let extractor = DefaultShardExtractor::new(16);
#[rustfmt::skip]
let cases: &[(&str, u64)] = &[
("entity-0", 15),
("entity-1", 12),
("entity-100", 12),
("entity-999", 10),
("a", 12),
("hello", 11),
("", 5),
("datum-sharding", 3),
("very-long-entity-id-that-exercises-fnv-loop-0123456789", 8),
];
for &(entity_id, expected) in cases {
let shard: String = ShardExtractor::<u8>::shard_id(&extractor, entity_id);
assert_eq!(
shard.parse::<u64>().unwrap(),
expected,
"shard mapping changed for {entity_id:?}: expected {expected}, got {shard}"
);
}
}
#[test]
fn push_bounded_round_caps_at_max_rounds() {
const MAX_ROUNDS: usize = 1024;
let new_round = |r: u64| RebalanceRound {
round: r,
reason: RebalanceReason::DeadOwner,
movements: vec![ShardMovement {
type_name: "test".to_owned(),
shard_id: format!("shard-{r}"),
from_node: "old".to_owned(),
to_node: "new".to_owned(),
generation: r,
}],
};
let mut queue = VecDeque::new();
assert_eq!(queue.len(), 0);
for r in 0..512 {
push_bounded_round(&mut queue, new_round(r));
}
assert_eq!(queue.len(), 512);
for r in 512..1024 {
push_bounded_round(&mut queue, new_round(r));
}
assert_eq!(queue.len(), 1024);
assert_eq!(queue.front().unwrap().round, 0);
for r in 1024..2050 {
push_bounded_round(&mut queue, new_round(r));
}
assert_eq!(queue.len(), 1024);
assert_eq!(queue.front().unwrap().round, 2050 - MAX_ROUNDS as u64);
assert_eq!(queue.back().unwrap().round, 2049);
for r in 2050..3050 {
push_bounded_round(&mut queue, new_round(r));
}
assert_eq!(queue.len(), 1024);
assert_eq!(queue.front().unwrap().round, 3050 - MAX_ROUNDS as u64);
assert_eq!(queue.back().unwrap().round, 3049);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn handoff_stops_live_entities_on_losing_owner() {
use datum_agent::{
ClusterAgent, ClusterAgentConfig, ClusterAgentHandle, NodeSessionConfig,
dcp::{DcpJobFactories, DcpServerConfig, DcpTcpServerConfig},
};
use datum_cluster::ClusterConfig;
use std::time::Duration;
const TYPE: &str = "handoff-entity-test";
#[derive(serde::Serialize, serde::Deserialize)]
enum Msg {
Ping,
}
struct TestEntity;
impl crate::EntityBehavior<Msg> for TestEntity {
fn handle(
&mut self,
_context: &crate::EntityContext,
_message: Msg,
) -> crate::ShardingResult<()> {
Ok(())
}
}
fn agent_config(node_id: &str, seeds: Vec<std::net::SocketAddr>) -> ClusterAgentConfig {
ClusterAgentConfig {
cluster: ClusterConfig {
node_id: node_id.to_owned(),
seed_nodes: seeds,
bind_addr: "127.0.0.1:0".parse().unwrap(),
advertise_addr: "127.0.0.1:0".parse().unwrap(),
gossip_interval: Duration::from_millis(60),
probe_timeout: Duration::from_millis(15),
suspect_timeout: Duration::from_secs(2),
downing_timeout: Duration::from_millis(150),
remove_down_after: Duration::from_secs(20),
event_buffer: 512,
..ClusterConfig::default()
},
dcp: DcpServerConfig {
node_id: node_id.to_owned(),
tcp: Some(DcpTcpServerConfig {
addr: "127.0.0.1:0".parse().unwrap(),
}),
..DcpServerConfig::default()
},
sessions: NodeSessionConfig {
request_timeout: Duration::from_secs(2),
command_buffer: 256,
reconnect_min_backoff: Duration::from_millis(20),
reconnect_max_backoff: Duration::from_millis(100),
..NodeSessionConfig::default()
},
..ClusterAgentConfig::default()
}
}
let mut agents = Vec::new();
for i in 0..2 {
let seeds: Vec<_> = agents
.first()
.map(|a: &ClusterAgentHandle| vec![a.cluster().advertise_addr()])
.unwrap_or_default();
agents.push(
ClusterAgent::start(
agent_config(&format!("hent-{i}"), seeds),
DcpJobFactories::new(),
)
.await
.expect("agent"),
);
}
tokio::time::sleep(Duration::from_millis(500)).await;
let sharding_config = ShardingConfig {
num_shards: 8,
rebalance_per_round: 8,
handoff_drain_delay: Duration::from_millis(10),
..ShardingConfig::default()
};
let mut handles = Vec::new();
for agent in &agents {
let handle = Sharding::init(agent, sharding_config.clone()).expect("sharding");
handle
.register_entity_type(TYPE, |_| TestEntity)
.await
.expect("register");
handles.push(handle);
}
for entity in 0..80 {
handles[0]
.entity_ref::<Msg>(TYPE, format!("e-{entity}"))
.tell(Msg::Ping)
.await
.expect("tell");
}
let table0 = handles[0].allocation_table(TYPE).await.expect("table");
let node0 = agents[0].cluster().node_id().to_owned();
let mut target_shard = None;
for entry in &table0.entries {
if entry.node_id == node0 {
let count = handles[0]
.test_live_entity_count(TYPE, &entry.shard_id)
.await
.expect("count");
if count > 0 {
target_shard = Some(entry.shard_id.clone());
break;
}
}
}
let target_shard =
target_shard.expect("node 0 must own at least one shard with live entities");
let before_count = handles[0]
.test_live_entity_count(TYPE, &target_shard)
.await
.expect("before count");
assert!(
before_count > 0,
"node 0 must have live entities for shard {target_shard} before rebalance"
);
let seed = agents[0].cluster().advertise_addr();
agents.push(
ClusterAgent::start(agent_config("hent-2", vec![seed]), DcpJobFactories::new())
.await
.expect("agent 3"),
);
let handle3 = Sharding::init(&agents[2], sharding_config).expect("sharding 3");
handle3
.register_entity_type(TYPE, |_| TestEntity)
.await
.expect("register 3");
handles.push(handle3);
tokio::time::sleep(Duration::from_millis(500)).await;
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
let mut moved = false;
while tokio::time::Instant::now() < deadline && !moved {
if let Ok(table) = handles[0].allocation_table(TYPE).await {
moved = table
.entries
.iter()
.any(|e| e.shard_id == target_shard && e.node_id != node0);
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
assert!(
moved,
"shard {target_shard} must move off node 0 after rebalance"
);
let after_count = handles[0]
.test_live_entity_count(TYPE, &target_shard)
.await
.expect("after count");
assert_eq!(
after_count, 0,
"node 0 must hold zero live entities for shard {target_shard} after handoff, got {after_count}"
);
for h in &handles {
h.shutdown().await;
}
for a in agents {
let _ = a.shutdown().await;
}
}
}