use crate::address::ensure_protocol;
use crate::errors::OxiaError;
use crate::errors::OxiaError::{UnexpectedStatus};
use crate::oxia::shard_assignment::ShardBoundaries;
use crate::oxia::{ShardAssignment, ShardAssignmentsRequest};
use crate::provider_manager::ProviderManager;
use backoff::{Error, ExponentialBackoff};
use dashmap::DashMap;
use log::{info, warn};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::{mpsc, Mutex};
use tokio::task::JoinHandle;
use tokio_util::sync::CancellationToken;
use tonic::codegen::tokio_stream::StreamExt;
pub struct ShardManagerOptions {
pub address: String,
pub namespace: String,
pub provider_manager: Arc<ProviderManager>,
}
pub(crate) struct Node {
pub service_address: String,
}
struct Inner {
provider_manager: Arc<ProviderManager>,
current_assignments: Arc<DashMap<i64, ShardAssignment>>,
}
pub struct ShardManager {
inner: Arc<Inner>,
context: CancellationToken,
assignment_handle: Mutex<Option<JoinHandle<()>>>,
}
impl Drop for ShardManager {
fn drop(&mut self) {
self.context.cancel();
}
}
async fn start_assignments_listener(
context: CancellationToken,
address: String,
namespace: String,
inner: Arc<Inner>,
init_tx: mpsc::Sender<()>,
) {
let op_defer = || {
let assignments_ref = inner.current_assignments.clone();
let local_address = address.clone();
let local_inner = inner.clone();
let local_context = context.clone();
let ns = namespace.clone();
let init_sender = init_tx.clone();
async move {
let mut provider = local_inner
.provider_manager
.get_provider(local_address)
.await
.map_err(|err| Error::transient(UnexpectedStatus(err.to_string())))?;
let mut streaming = provider
.get_shard_assignments(ShardAssignmentsRequest {
namespace: ns.clone(),
})
.await
.map_err(|err| Error::transient(UnexpectedStatus(err.to_string())))?
.into_inner();
loop {
tokio::select! {
_ = local_context.cancelled() => {
info!("Close shards assignment stream due to context canceled.");
return Ok(());
}
result = streaming.next() => {
if result.is_none() {
info!("Close shards assignment stream due to stream closed.");
return Ok(());
}
match result.unwrap() {
Ok(assignments) => {
let ns_assignments = assignments.namespaces.get(&ns);
if ns_assignments.is_none() {
continue;
}
assignments_ref.clear();
for sa in &ns_assignments.unwrap().assignments {
assignments_ref.insert(sa.shard, sa.clone());
}
if !init_sender.is_closed() {
let _ = init_sender.send(()).await;
}
}
Err(stream_status) => {
return Err(Error::transient(UnexpectedStatus(stream_status.to_string())));
}
}
}
}
}
}
};
let backoff = ExponentialBackoff::default();
let _ = backoff::future::retry_notify(backoff, op_defer, |err, duration| {
warn!(
"Transient failure receiving shard assignments. error: {:?} retry-after: {:?}.",
err, duration
)
})
.await;
}
impl ShardManager {
pub async fn new(options: ShardManagerOptions) -> Result<Self, OxiaError> {
let context = CancellationToken::new();
let inner = Arc::new(Inner {
provider_manager: options.provider_manager,
current_assignments: Arc::new(DashMap::new()),
});
let (init_tx, mut init_rx) = mpsc::channel(1);
let assignment_handle = tokio::spawn(start_assignments_listener(
context.clone(),
options.address,
options.namespace.clone(),
inner.clone(),
init_tx,
));
init_rx.recv().await;
let sm = ShardManager {
inner,
context,
assignment_handle: Mutex::new(Some(assignment_handle)),
};
Ok(sm)
}
pub async fn shutdown(self) -> Result<(), OxiaError> {
self.context.cancel();
let mut handle_guard = self.assignment_handle.lock().await;
if let Some(handle) = handle_guard.take() {
handle
.await
.map_err(|err| UnexpectedStatus(err.to_string()))?;
}
Ok(())
}
pub fn get_leader(&self, shard_id: i64) -> Option<Node> {
let assignment_option = self.inner.current_assignments.get(&shard_id);
assignment_option.map(|v| Node {
service_address: ensure_protocol(v.leader.clone()),
})
}
pub fn get_shard(&self, key: &str) -> Option<i64> {
let code = xxhash_rust::xxh32::xxh32(key.as_bytes(), 0);
for entry in self.inner.current_assignments.iter() {
match entry.shard_boundaries.unwrap() {
ShardBoundaries::Int32HashRange(range) => {
if range.min_hash_inclusive <= code && code <= range.max_hash_inclusive {
return Some(entry.shard);
}
}
}
}
None
}
pub fn get_shards_leader(&self) -> HashMap<i64, Node> {
let mut map = HashMap::with_capacity(self.inner.current_assignments.len());
for item in self.inner.current_assignments.iter() {
map.insert(
item.shard,
Node {
service_address: ensure_protocol(item.leader.clone()),
},
);
}
map
}
pub fn get_shard_leader(&self, key: &str) -> Option<Node> {
let code = xxhash_rust::xxh32::xxh32(key.as_bytes(), 0);
for entry in self.inner.current_assignments.iter() {
match entry.shard_boundaries.unwrap() {
ShardBoundaries::Int32HashRange(range) => {
if range.min_hash_inclusive <= code && code <= range.max_hash_inclusive {
return Some(Node {
service_address: ensure_protocol(entry.leader.clone()),
});
}
}
}
}
None
}
}