liboxia 0.0.2

Liboxia is a Rust library designed for both native Rust applications and for integration with other languages via its C Foreign Function Interface (FFI). It serves as a client SDK for Oxia, a distributed key-value store, enabling robust, asynchronous data operations.
Documentation
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
    }

    // todo: consider iterator to avoid clone
    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
    }
}