use crate::bucket_index::BucketIndex;
use crate::bucket_locks::BucketLocks;
use crate::indexer::Indexer;
use crate::ingest_counter;
use crate::quantizer::Quantizer;
use crate::router::{Router, RoutingTable};
use crate::storage::{Bucket, BucketMeta, BucketMetaStatus, Storage, Vector};
use crate::vector_index::VectorIndex;
use crate::wal::runtime::Walrus;
use anyhow::Result;
use futures::executor::block_on;
use log::{debug, error, warn};
use parking_lot::RwLock;
use std::collections::{HashMap, HashSet};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::thread;
use std::time::Duration;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RebalanceTaskKind {
Split,
Merge,
Rebalance,
}
type RebalanceFailHook = Arc<dyn Fn(RebalanceTaskKind) -> bool + Send + Sync>;
static FAIL_HOOK: StdMutex<Option<RebalanceFailHook>> = StdMutex::new(None);
#[derive(Debug)]
pub struct DeleteCommand {
pub vector_id: u64,
pub bucket_hint: Option<u64>,
pub respond_to: futures::channel::oneshot::Sender<anyhow::Result<()>>,
}
pub fn set_rebalance_fail_hook<F>(hook: F)
where
F: Fn(RebalanceTaskKind) -> bool + Send + Sync + 'static,
{
*FAIL_HOOK.lock().expect("fail hook poisoned") = Some(Arc::new(hook));
}
pub fn clear_rebalance_fail_hook() {
*FAIL_HOOK.lock().expect("fail hook poisoned") = None;
}
fn should_fail(kind: RebalanceTaskKind) -> bool {
if let Some(h) = FAIL_HOOK.lock().expect("fail hook poisoned").as_ref() {
return h(kind);
}
let _ = kind;
false
}
pub(crate) struct RebalanceState {
storage: Storage,
vector_index: Arc<VectorIndex>,
bucket_index: Arc<BucketIndex>,
wal: Arc<Walrus>,
routing: Arc<RoutingTable>,
centroids: RwLock<HashMap<u64, Vec<f32>>>,
bucket_sizes: RwLock<HashMap<u64, usize>>,
retired: RwLock<HashSet<u64>>,
next_bucket_id: AtomicU64,
bucket_locks: Arc<BucketLocks>,
}
impl RebalanceState {
fn new(
storage: Storage,
vector_index: Arc<VectorIndex>,
bucket_index: Arc<BucketIndex>,
routing: Arc<RoutingTable>,
bucket_locks: Arc<BucketLocks>,
) -> Self {
Self {
wal: storage.wal.clone(),
storage,
vector_index,
bucket_index,
routing,
centroids: RwLock::new(HashMap::new()),
bucket_sizes: RwLock::new(HashMap::new()),
retired: RwLock::new(HashSet::new()),
next_bucket_id: AtomicU64::new(0),
bucket_locks,
}
}
fn prime_centroids(&self, buckets: &[Bucket]) {
let mut map = self.centroids.write();
let mut sizes = self.bucket_sizes.write();
let mut max_id = 0;
for b in buckets {
map.insert(b.id, b.centroid.clone());
sizes.insert(b.id, b.vectors.len());
if b.id > max_id {
max_id = b.id;
}
}
self.next_bucket_id.store(max_id + 1, Ordering::Release);
}
fn allocate_bucket_id(&self) -> u64 {
self.next_bucket_id.fetch_add(1, Ordering::AcqRel)
}
fn refresh_sizes(&self) -> HashMap<u64, usize> {
let ids: Vec<u64> = self.centroids.read().keys().cloned().collect();
let prev_sizes = self.bucket_sizes.read().clone();
let mut fresh = HashMap::new();
for id in ids {
let topic = crate::storage::Storage::topic_for(id);
let count = self.wal.get_topic_entry_count(&topic) as usize;
let stabilized = count.max(prev_sizes.get(&id).cloned().unwrap_or(0));
fresh.insert(id, stabilized);
}
let mut sizes = self.bucket_sizes.write();
sizes.clear();
sizes.extend(fresh.iter().map(|(k, v)| (*k, *v)));
let sum: u64 = fresh.values().map(|v| *v as u64).sum();
let inserted = ingest_counter::get();
if sum < inserted {
debug!(
"rebalance: wal counts sum {} is less than total inserted {}; proceeding with monotonic sizes",
sum, inserted
);
}
sizes.clone()
}
fn lock_for(&self, bucket_id: u64) -> Arc<futures::lock::Mutex<()>> {
self.bucket_locks.lock_for(bucket_id)
}
pub(crate) fn load_bucket_vectors(&self, bucket_id: u64) -> Option<Vec<Vector>> {
let chunks = Storage::get_chunks_sync(self.storage.wal.clone(), bucket_id).ok()?;
let mut vector_map: HashMap<u64, Vector> = HashMap::new();
for chunk in chunks {
if chunk.len() < 8 {
continue;
}
let mut len_bytes = [0u8; 8];
len_bytes.copy_from_slice(&chunk[0..8]);
let archive_len = u64::from_le_bytes(len_bytes) as usize;
if 8 + archive_len > chunk.len() || archive_len < 16 {
continue;
}
let mut off = 8;
let mut id_bytes = [0u8; 8];
id_bytes.copy_from_slice(&chunk[off..off + 8]);
off += 8;
let mut dim_bytes = [0u8; 8];
dim_bytes.copy_from_slice(&chunk[off..off + 8]);
off += 8;
let dim = u64::from_le_bytes(dim_bytes) as usize;
let Some(expected_bytes) = dim.checked_mul(4) else {
continue;
};
if off + expected_bytes > chunk.len() {
continue;
}
let mut data = Vec::with_capacity(dim);
let data_bytes = &chunk[off..off + expected_bytes];
for chunked in data_bytes.chunks_exact(4) {
let mut fb = [0u8; 4];
fb.copy_from_slice(chunked);
data.push(f32::from_bits(u32::from_le_bytes(fb)));
}
let id = u64::from_le_bytes(id_bytes);
if data.is_empty() {
vector_map.remove(&id);
} else {
vector_map.insert(id, Vector { id, data });
}
}
if vector_map.is_empty() {
None
} else {
Some(vector_map.into_values().collect())
}
}
fn retire_bucket_local(&self, bucket_id: u64) {
self.centroids.write().remove(&bucket_id);
self.bucket_sizes.write().remove(&bucket_id);
self.retired.write().insert(bucket_id);
}
fn retire_bucket_io(&self, bucket_id: u64) {
let _ = block_on(self.storage.put_bucket_meta(&BucketMeta {
bucket_id,
status: BucketMetaStatus::Retired,
}));
}
fn mark_bucket_checkpointed(&self, bucket_id: u64) {
let topic = crate::storage::Storage::topic_for(bucket_id);
let max_bytes = 16 * 1024 * 1024;
loop {
match self.wal.batch_read_for_topic(&topic, max_bytes, true, None) {
Ok(entries) => {
if entries.is_empty() {
break;
}
}
Err(e) => {
warn!("rebalance: checkpoint drain failed for {}: {:?}", topic, e);
break;
}
}
}
}
fn rebuild_router(&self, changed_buckets: Vec<u64>) {
let centroids_map = self.centroids.read();
let bucket_count = centroids_map.len();
if centroids_map.is_empty() {
return;
}
let mut centroids: Vec<(u64, Vec<f32>)> = Vec::with_capacity(bucket_count);
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for (id, c) in centroids_map.iter() {
for &val in c {
if val < min {
min = val;
}
if val > max {
max = val;
}
}
centroids.push((*id, c.clone()));
}
drop(centroids_map);
let Some((min, max)) = Quantizer::compute_bounds_from_minmax(min, max) else {
return;
};
let quantizer = Quantizer::new(min, max);
let mut router = Router::new(100_000, quantizer);
for (id, centroid) in ¢roids {
router.add_centroid(*id, centroid);
}
let version = self.routing.install(router, changed_buckets);
if log::log_enabled!(log::Level::Debug) {
let sizes_map = self.bucket_sizes.read();
let mut sizes: Vec<usize> = sizes_map.values().copied().collect();
sizes.sort_unstable();
debug!(
"rebalance: published router version {} (buckets={}, sizes={:?})",
version, bucket_count, sizes
);
}
}
fn handle_split_sync(&self, bucket_id: u64) {
if should_fail(RebalanceTaskKind::Split) {
debug!(
"rebalance: injected failure for split on bucket {}",
bucket_id
);
return;
}
let lock = self.lock_for(bucket_id);
let _guard = block_on(lock.lock());
if self.retired.read().contains(&bucket_id) {
debug!(
"rebalance: split skipped, bucket {} already retired",
bucket_id
);
return;
}
let vectors = match self.load_bucket_vectors(bucket_id) {
Some(v) => v,
None => {
log::debug!("rebalance: split skipped, bucket {} not found", bucket_id);
return;
}
};
let mut bucket = Bucket::new(bucket_id, Vec::new());
bucket.vectors = vectors;
let splits = Indexer::split_bucket_once(bucket);
if splits.is_empty() {
return;
}
let mut new_entries = Vec::new();
let mut bucket_index_updates = Vec::new();
for mut split in splits {
let new_id = self.allocate_bucket_id();
split.id = new_id;
if let Err(e) = block_on(self.storage.put_chunk(&split)) {
error!(
"rebalance: failed to persist split bucket {} -> {}: {:?}",
bucket_id, new_id, e
);
continue;
}
let _ = block_on(self.storage.put_bucket_meta(&BucketMeta {
bucket_id: new_id,
status: BucketMetaStatus::Active,
}));
let ids: Vec<u64> = split.vectors.iter().map(|v| v.id).collect();
bucket_index_updates.push((new_id, ids));
new_entries.push((new_id, split.centroid, split.vectors.len()));
}
if new_entries.is_empty() {
return;
}
for (id, ids) in bucket_index_updates {
if let Err(e) = self.bucket_index.put_batch(id, &ids) {
error!(
"rebalance: failed to update bucket index for split bucket {}: {:?}",
id, e
);
}
}
self.retire_bucket_local(bucket_id);
self.retire_bucket_io(bucket_id);
let mut centroids = self.centroids.write();
let mut sizes = self.bucket_sizes.write();
let new_count = new_entries.len();
let mut changed = Vec::with_capacity(1 + new_count);
changed.push(bucket_id);
for (id, centroid, size) in new_entries {
centroids.insert(id, centroid);
sizes.insert(id, size);
changed.push(id);
}
drop(sizes);
drop(centroids);
self.rebuild_router(changed);
self.mark_bucket_checkpointed(bucket_id);
log::info!(
"rebalance: split {} into {} buckets (new_total={})",
bucket_id,
new_count,
self.centroids.read().len()
);
}
}
pub struct RebalanceWorker {
state: Arc<RebalanceState>,
pub delete_tx: async_channel::Sender<DeleteCommand>,
}
impl RebalanceWorker {
pub fn new_for_tests(
storage: Storage,
vector_index: Arc<VectorIndex>,
bucket_index: Arc<BucketIndex>,
routing: Arc<RoutingTable>,
bucket_locks: Arc<BucketLocks>,
) -> Self {
let state = Arc::new(RebalanceState::new(
storage,
vector_index,
bucket_index,
routing,
bucket_locks,
));
let (delete_tx, _delete_rx) = async_channel::bounded(1024);
Self { state, delete_tx }
}
pub fn spawn(
storage: Storage,
vector_index: Arc<VectorIndex>,
bucket_index: Arc<BucketIndex>,
routing: Arc<RoutingTable>,
pin_cpu: Option<usize>,
bucket_locks: Arc<BucketLocks>,
) -> Self {
let state = Arc::new(RebalanceState::new(
storage,
vector_index,
bucket_index,
routing,
bucket_locks,
));
let state_clone = state.clone();
let (delete_tx, delete_rx) = async_channel::bounded(1024);
let name = "rebalance-loop".to_string();
match pin_cpu {
Some(cpu) => {
let builder =
glommio::LocalExecutorBuilder::new(glommio::Placement::Fixed(cpu)).name(&name);
std::thread::spawn(move || {
builder
.make()
.expect("failed to create rebalance executor")
.run(run_autonomous_loop(state_clone, delete_rx));
});
}
None => {
thread::Builder::new()
.name(name.clone())
.spawn(move || {
glommio::LocalExecutorBuilder::default()
.name(&name)
.make()
.expect("failed to create default rebalance executor")
.run(run_autonomous_loop(state_clone, delete_rx));
})
.expect("rebalance worker");
}
}
Self { state, delete_tx }
}
pub fn spawn_delete_only(
storage: Storage,
vector_index: Arc<VectorIndex>,
bucket_index: Arc<BucketIndex>,
routing: Arc<RoutingTable>,
bucket_locks: Arc<BucketLocks>,
) -> (Self, std::thread::JoinHandle<()>) {
let state = Arc::new(RebalanceState::new(
storage,
vector_index,
bucket_index,
routing,
bucket_locks,
));
let state_clone = state.clone();
let (delete_tx, delete_rx) = async_channel::bounded(1024);
let name = "delete-worker".to_string();
let handle = thread::Builder::new()
.name(name.clone())
.spawn(move || {
glommio::LocalExecutorBuilder::default()
.name(&name)
.make()
.expect("failed to create delete executor")
.run(run_delete_loop(state_clone, delete_rx));
})
.expect("delete worker");
(Self { state, delete_tx }, handle)
}
pub async fn prime_centroids(&self, buckets: &[Bucket]) -> Result<()> {
self.state.prime_centroids(buckets);
self.state.rebuild_router(Vec::new());
for b in buckets {
self.state
.storage
.put_bucket_meta(&BucketMeta {
bucket_id: b.id,
status: BucketMetaStatus::Active,
})
.await?;
}
Ok(())
}
pub fn snapshot_sizes(&self) -> HashMap<u64, usize> {
self.state.refresh_sizes()
}
pub fn close(&self) {}
pub async fn delete(&self, vector_id: u64, bucket_hint: Option<u64>) -> anyhow::Result<()> {
let (tx, rx) = futures::channel::oneshot::channel();
self.delete_tx
.send(DeleteCommand {
vector_id,
bucket_hint,
respond_to: tx,
})
.await
.map_err(|_| anyhow::anyhow!("rebalance delete queue closed"))?;
rx.await
.map_err(|e| anyhow::anyhow!("rebalance delete canceled: {:?}", e))?
}
pub fn delete_inline_blocking(
&self,
vector_id: u64,
bucket_hint: Option<u64>,
) -> anyhow::Result<()> {
let (tx, rx) = futures::channel::oneshot::channel();
let cmd = DeleteCommand {
vector_id,
bucket_hint,
respond_to: tx,
};
futures::executor::block_on(handle_delete(&self.state, cmd))?;
futures::executor::block_on(rx)
.map_err(|e| anyhow::anyhow!("rebalance delete canceled: {:?}", e))?
}
}
impl Clone for RebalanceWorker {
fn clone(&self) -> Self {
Self {
state: self.state.clone(),
delete_tx: self.delete_tx.clone(),
}
}
}
async fn run_delete_loop(
state: Arc<RebalanceState>,
delete_rx: async_channel::Receiver<DeleteCommand>,
) {
while let Ok(cmd) = delete_rx.recv().await {
if let Err(e) = handle_delete(&state, cmd).await {
warn!("rebalance: delete failed: {:?}", e);
}
}
}
async fn run_autonomous_loop(
state: Arc<RebalanceState>,
delete_rx: async_channel::Receiver<DeleteCommand>,
) {
let threshold: usize = std::env::var("SATORI_REBALANCE_THRESHOLD")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(2000);
loop {
while let Ok(cmd) = delete_rx.try_recv() {
if let Err(e) = handle_delete(&state, cmd).await {
warn!("rebalance: delete failed: {:?}", e);
}
}
let sizes = state.refresh_sizes();
let mut max_id = 0;
let mut max_size = 0;
for (id, size) in sizes {
if size > max_size {
max_size = size;
max_id = id;
}
}
if max_size > threshold {
let state_ref = state.clone();
let _ = glommio::executor()
.spawn_blocking(move || {
state_ref.handle_split_sync(max_id);
})
.await;
} else {
glommio::timer::Timer::new(Duration::from_millis(500)).await;
}
}
}
async fn perform_delete(
state: &Arc<RebalanceState>,
vector_id: u64,
bucket_hint: Option<u64>,
) -> anyhow::Result<()> {
let bucket_id = if let Some(b) = bucket_hint {
b
} else {
let found = state
.bucket_index
.get_many(&[vector_id])
.map_err(|e| anyhow::anyhow!("bucket index lookup failed: {:?}", e))?;
match found.first() {
Some((_, b)) => *b,
None => {
return Ok(());
}
}
};
let lock = state.lock_for(bucket_id);
let _guard = lock.lock().await;
let vectors = match state.load_bucket_vectors(bucket_id) {
Some(v) => v,
None => {
return Ok(());
}
};
let mut remaining: Vec<Vector> = Vec::with_capacity(vectors.len());
let mut removed = false;
for v in vectors {
if v.id == vector_id {
removed = true;
} else {
remaining.push(v);
}
}
if !removed {
return Ok(());
}
let _centroid = state
.centroids
.read()
.get(&bucket_id)
.cloned()
.unwrap_or_default();
let topic = crate::storage::Storage::topic_for(bucket_id);
let entries_before = state.storage.wal.get_topic_entry_count(&topic);
Storage::put_chunk_raw_sync(state.storage.wal.clone(), bucket_id, &topic, &remaining)
.map_err(|e| anyhow::anyhow!("rewrite bucket after delete failed: {:?}", e))?;
let tombstone = Vector::new(vector_id, Vec::new());
if let Err(e) =
Storage::put_chunk_raw_sync(state.storage.wal.clone(), bucket_id, &topic, &[tombstone])
{
warn!(
"rebalance: failed to append delete tombstone for {}: {:?}",
vector_id, e
);
}
let ids = [vector_id];
if let Err(e) = state.vector_index.delete_batch(&ids) {
warn!(
"rebalance: delete failed from vector index for {}: {:?}",
vector_id, e
);
}
if let Err(e) = state.bucket_index.delete_batch(&ids) {
warn!(
"rebalance: delete failed from bucket index for {}: {:?}",
vector_id, e
);
}
let mut remaining = entries_before;
while remaining > 0 {
const MIN_ENTRY_BYTES: usize = 24; let max_entries = remaining.min(2000) as usize;
let max_bytes = max_entries
.saturating_mul(MIN_ENTRY_BYTES)
.max(MIN_ENTRY_BYTES);
match state
.storage
.wal
.batch_read_for_topic(&topic, max_bytes, true, None)
{
Ok(batch) => {
if batch.is_empty() {
break;
}
let consumed = batch.len().min(max_entries) as u64;
remaining = remaining.saturating_sub(consumed);
if consumed == 0 {
break;
}
}
Err(e) => {
warn!("rebalance: checkpoint drain failed for {}: {:?}", topic, e);
break;
}
}
}
Ok(())
}
async fn handle_delete(state: &Arc<RebalanceState>, cmd: DeleteCommand) -> anyhow::Result<()> {
let result = perform_delete(state, cmd.vector_id, cmd.bucket_hint).await;
let send_payload = result
.as_ref()
.map(|_| ())
.map_err(|e| anyhow::anyhow!("{:?}", e));
let _ = cmd.respond_to.send(send_payload);
result
}
#[cfg(test)]
mod tests {
use super::*;
pub(crate) fn compute_centroid(vectors: &[Vector]) -> Vec<f32> {
if vectors.is_empty() {
return Vec::new();
}
let dim = vectors[0].data.len();
let mut sums = vec![0.0f32; dim];
for v in vectors {
for (i, val) in v.data.iter().enumerate() {
sums[i] += *val;
}
}
let count = vectors.len() as f32;
for s in sums.iter_mut() {
*s /= count;
}
sums
}
#[test]
fn centroid_of_vectors_is_mean() {
let vectors = vec![
Vector::new(0, vec![1.0, 3.0]),
Vector::new(1, vec![3.0, 5.0]),
];
let centroid = compute_centroid(&vectors);
assert_eq!(centroid, vec![2.0, 4.0]);
}
#[test]
fn centroid_of_empty_is_empty() {
let centroid = compute_centroid(&[]);
assert!(centroid.is_empty());
}
#[test]
fn centroid_of_single_is_identity() {
let vectors = vec![Vector::new(0, vec![5.0, 10.0, 15.0])];
let centroid = compute_centroid(&vectors);
assert_eq!(centroid, vec![5.0, 10.0, 15.0]);
}
#[test]
fn centroid_many_vectors() {
let n = 1000;
let vectors: Vec<Vector> = (0..n)
.map(|i| Vector::new(i as u64, vec![i as f32, (i * 2) as f32]))
.collect();
let centroid = compute_centroid(&vectors);
let expected_x = (0..n).sum::<usize>() as f32 / n as f32;
let expected_y = (0..n).map(|i| (i * 2) as f32).sum::<f32>() / n as f32;
assert!((centroid[0] - expected_x).abs() < 0.01);
assert!((centroid[1] - expected_y).abs() < 0.01);
}
#[test]
fn centroid_negative_values() {
let vectors = vec![
Vector::new(0, vec![-10.0, -20.0]),
Vector::new(1, vec![10.0, 20.0]),
];
let centroid = compute_centroid(&vectors);
assert_eq!(centroid, vec![0.0, 0.0]);
}
#[test]
fn centroid_zero_dimension() {
let vectors = vec![Vector::new(0, vec![]), Vector::new(1, vec![])];
let centroid = compute_centroid(&vectors);
assert!(centroid.is_empty());
}
}