use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Weak};
use std::thread::JoinHandle;
use std::time::Duration;
use super::YantrikDB;
const DRAIN_BATCH_SIZE: usize = 64;
const IDLE_POLL_INTERVAL: Duration = Duration::from_millis(100);
const ERROR_BACKOFF_INTERVAL: Duration = Duration::from_millis(100);
pub struct MaterializerGuard {
shutdown: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
impl Drop for MaterializerGuard {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Relaxed);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
pub fn spawn_materializers(db: &Arc<YantrikDB>, count: usize) -> Vec<MaterializerGuard> {
let mut guards = Vec::with_capacity(count);
for worker_id in 0..count {
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_clone = Arc::clone(&shutdown);
let weak = Arc::downgrade(db);
let handle = std::thread::Builder::new()
.name(format!("yantrikdb-materializer-{worker_id}"))
.spawn(move || worker_loop(weak, shutdown_clone, worker_id))
.expect("spawn materializer thread");
guards.push(MaterializerGuard {
shutdown,
handle: Some(handle),
});
}
guards
}
fn worker_loop(weak: Weak<YantrikDB>, shutdown: Arc<AtomicBool>, worker_id: usize) {
tracing::debug!(worker_id, "materializer worker started");
while !shutdown.load(Ordering::Relaxed) {
let Some(db) = weak.upgrade() else {
tracing::debug!(worker_id, "engine dropped — materializer exiting");
break;
};
match db.apply_pending_ops_once(DRAIN_BATCH_SIZE) {
Ok(0) => {
drop(db);
std::thread::sleep(IDLE_POLL_INTERVAL);
}
Ok(n) => {
tracing::trace!(worker_id, applied = n, "drained batch");
drop(db);
}
Err(e) => {
tracing::warn!(worker_id, error = %e, "drain failed; backing off");
drop(db);
std::thread::sleep(ERROR_BACKOFF_INTERVAL);
}
}
}
tracing::debug!(worker_id, "materializer worker exited");
}
pub fn recommended_worker_count() -> usize {
std::thread::available_parallelism()
.map(|n| n.get() / 2)
.unwrap_or(2)
.clamp(2, 16)
}
pub struct CompactorGuard {
shutdown: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
impl Drop for CompactorGuard {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Relaxed);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
const COMPACTOR_INTERVAL: Duration = Duration::from_millis(250);
pub fn spawn_compactor(db: &Arc<YantrikDB>) -> CompactorGuard {
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_clone = Arc::clone(&shutdown);
let weak = Arc::downgrade(db);
let handle = std::thread::Builder::new()
.name("yantrikdb-compactor".to_string())
.spawn(move || compactor_loop(weak, shutdown_clone))
.expect("spawn compactor thread");
CompactorGuard {
shutdown,
handle: Some(handle),
}
}
fn compactor_loop(weak: Weak<YantrikDB>, shutdown: Arc<AtomicBool>) {
tracing::debug!("compactor started");
while !shutdown.load(Ordering::Relaxed) {
let Some(db) = weak.upgrade() else {
tracing::debug!("engine dropped — compactor exiting");
break;
};
if db.vec_index.should_compact() {
match db.vec_index.compact() {
Ok(0) => {}
Ok(n) => {
tracing::debug!(applied = n, "compaction drained delta into cold");
}
Err(e) => {
tracing::warn!(error = %e, "compaction failed; retrying next tick");
}
}
}
drop(db);
std::thread::sleep(COMPACTOR_INTERVAL);
}
tracing::debug!("compactor exited");
}
#[cfg(test)]
mod tests {
use super::*;
use crate::YantrikDB;
fn open_test_db() -> Arc<YantrikDB> {
Arc::new(YantrikDB::new(":memory:", 64).expect("open"))
}
#[test]
fn worker_drains_pending_ops() {
let db = open_test_db();
for i in 0..5 {
db.log_op_pending(
"record",
Some(&format!("rid_{i}")),
&serde_json::json!({}),
None,
None,
)
.unwrap();
}
assert_eq!(db.count_pending_ops().unwrap(), 5);
let _guards = spawn_materializers(&db, 1);
let mut tries = 0;
while db.count_pending_ops().unwrap() > 0 && tries < 20 {
std::thread::sleep(Duration::from_millis(50));
tries += 1;
}
assert_eq!(
db.count_pending_ops().unwrap(),
0,
"worker should drain all 5 pending ops within 1s"
);
}
#[test]
fn guard_drop_shuts_down_worker() {
let db = open_test_db();
let guards = spawn_materializers(&db, 2);
for i in 0..3 {
db.log_op_pending(
"record",
Some(&format!("rid_{i}")),
&serde_json::json!({}),
None,
None,
)
.unwrap();
}
std::thread::sleep(Duration::from_millis(200));
let start = std::time::Instant::now();
drop(guards);
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(1),
"guard drop should join workers within 1s, took {elapsed:?}"
);
}
#[test]
fn engine_drop_lets_workers_exit_via_weak_upgrade_fail() {
let db = open_test_db();
let guards = spawn_materializers(&db, 1);
drop(db);
let start = std::time::Instant::now();
drop(guards);
let elapsed = start.elapsed();
assert!(
elapsed < Duration::from_secs(1),
"engine drop should let worker exit within 1s, took {elapsed:?}"
);
}
#[test]
fn multiple_workers_dont_double_apply() {
let db = open_test_db();
for i in 0..50 {
db.log_op_pending(
"record",
Some(&format!("rid_{i}")),
&serde_json::json!({}),
None,
None,
)
.unwrap();
}
let _guards = spawn_materializers(&db, 4);
let mut tries = 0;
while db.count_pending_ops().unwrap() > 0 && tries < 40 {
std::thread::sleep(Duration::from_millis(50));
tries += 1;
}
assert_eq!(db.count_pending_ops().unwrap(), 0);
let conn = db.read_conn();
let total: i64 = conn
.query_row(
"SELECT COUNT(*) FROM oplog WHERE applied = 1 AND op_type = 'record'",
[],
|row| row.get(0),
)
.unwrap();
assert_eq!(total, 50);
}
#[test]
fn recommended_worker_count_in_range() {
let n = recommended_worker_count();
assert!((2..=16).contains(&n), "expected [2,16], got {n}");
}
#[test]
fn compactor_drains_delta_periodically() {
let db = open_test_db();
let _guard = spawn_compactor(&db);
for i in 0..250 {
let emb: Vec<f32> = (0..64).map(|j| (i + j) as f32 * 0.001).collect();
let norm: f32 = emb.iter().map(|x| x * x).sum::<f32>().sqrt();
let normalized: Vec<f32> = emb.iter().map(|x| x / norm).collect();
let seq = db.vec_seq.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
db.vec_index.append(format!("rid_{i}"), normalized, seq).unwrap();
}
assert!(
db.vec_index.delta_len() >= 128,
"delta should have crossed compact threshold (half of 256), got {}",
db.vec_index.delta_len()
);
let mut tries = 0;
while db.vec_index.cold_len() < 200 && tries < 40 {
std::thread::sleep(Duration::from_millis(100));
tries += 1;
}
assert!(
db.vec_index.cold_len() >= 200,
"compactor should have moved >=200 entries to cold within 4s, got cold={} delta={}",
db.vec_index.cold_len(),
db.vec_index.delta_len()
);
assert!(
db.vec_index.delta_len() < 128,
"delta should be drained below half-cap, got {}",
db.vec_index.delta_len()
);
}
#[test]
fn compactor_guard_drop_shuts_down_clean() {
let db = open_test_db();
let guard = spawn_compactor(&db);
std::thread::sleep(Duration::from_millis(100));
let start = std::time::Instant::now();
drop(guard);
assert!(
start.elapsed() < Duration::from_secs(2),
"compactor guard drop must join within 2s"
);
}
}