use std::collections::{BTreeSet, HashMap};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use tokio::sync::{Notify, watch};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IndexState {
Ready,
Building {
done: usize,
total: usize,
elapsed: Duration,
},
Stale { pending: BTreeSet<String> },
}
impl IndexState {
pub fn is_building(&self) -> bool {
matches!(self, IndexState::Building { .. })
}
pub fn is_indexing(&self) -> bool {
!matches!(self, IndexState::Ready)
}
pub fn percent(&self) -> Option<u32> {
match self {
IndexState::Building { done, total, .. } if *total > 0 => {
Some(((*done * 100) / *total).min(99) as u32)
}
IndexState::Building { .. } => Some(0),
_ => None,
}
}
pub fn eta_secs(&self) -> Option<u64> {
match self {
IndexState::Building {
done,
total,
elapsed,
} if *done > 0 && *total > *done => {
let per_unit = elapsed.as_secs_f64() / *done as f64;
Some((per_unit * (*total - *done) as f64).ceil().max(1.0) as u64)
}
IndexState::Building { .. } => Some(5),
_ => None,
}
}
}
struct Entry {
tx: watch::Sender<IndexState>,
done: AtomicUsize,
total: AtomicUsize,
started: Mutex<Instant>,
}
impl Entry {
fn new() -> Self {
Self {
tx: watch::channel(IndexState::Ready).0,
done: AtomicUsize::new(0),
total: AtomicUsize::new(0),
started: Mutex::new(Instant::now()),
}
}
fn started(&self) -> Instant {
*self.started.lock().unwrap_or_else(|e| e.into_inner())
}
}
#[derive(Default)]
pub struct IndexStates {
entries: Mutex<HashMap<u64, Arc<Entry>>>,
watcher_inflight: AtomicUsize,
watcher_idle: Notify,
}
impl IndexStates {
fn entry(&self, project_id: u64) -> Arc<Entry> {
let mut entries = self.entries.lock().unwrap_or_else(|e| e.into_inner());
Arc::clone(
entries
.entry(project_id)
.or_insert_with(|| Arc::new(Entry::new())),
)
}
pub fn get(&self, project_id: u64) -> IndexState {
let entry = self.entry(project_id);
let state = entry.tx.borrow().clone();
match state {
IndexState::Building { .. } => IndexState::Building {
done: entry.done.load(Ordering::Relaxed),
total: entry.total.load(Ordering::Relaxed),
elapsed: entry.started().elapsed(),
},
other => other,
}
}
pub fn mark_building(&self, project_id: u64) {
let entry = self.entry(project_id);
if entry.tx.borrow().is_building() {
return;
}
entry.done.store(0, Ordering::Relaxed);
entry.total.store(0, Ordering::Relaxed);
*entry.started.lock().unwrap_or_else(|e| e.into_inner()) = Instant::now();
entry.tx.send_replace(IndexState::Building {
done: 0,
total: 0,
elapsed: Duration::ZERO,
});
}
pub fn set_total(&self, project_id: u64, total: usize) {
let entry = self.entry(project_id);
entry.total.store(total, Ordering::Relaxed);
}
pub fn add_progress(&self, project_id: u64, n: usize) {
self.entry(project_id).done.fetch_add(n, Ordering::Relaxed);
}
pub fn mark_stale(&self, project_id: u64, pending: BTreeSet<String>) {
let entry = self.entry(project_id);
if entry.tx.borrow().is_building() {
return;
}
entry.tx.send_replace(IndexState::Stale { pending });
}
pub fn file_done(&self, project_id: u64, path: &str) {
self.entry(project_id).tx.send_modify(|state| {
if let IndexState::Stale { pending } = state {
pending.remove(path);
}
});
}
pub fn mark_ready(&self, project_id: u64) {
self.entry(project_id).tx.send_replace(IndexState::Ready);
}
pub async fn wait_until_ready(&self, project_id: u64, limit: Duration) -> IndexState {
let mut rx = self.entry(project_id).tx.subscribe();
let _ = tokio::time::timeout(limit, rx.wait_for(|state| !state.is_building())).await;
self.get(project_id)
}
pub fn ready_on_drop(self: &Arc<Self>, project_id: u64) -> ReadyOnDrop {
ReadyOnDrop {
states: Arc::clone(self),
project_id,
}
}
pub fn watcher_begin(&self) {
self.watcher_inflight.fetch_add(1, Ordering::SeqCst);
}
pub fn watcher_end(&self) {
if self.watcher_inflight.fetch_sub(1, Ordering::SeqCst) == 1 {
self.watcher_idle.notify_waiters();
}
}
pub async fn wait_watcher_idle(&self, limit: Duration) {
let _ = tokio::time::timeout(limit, async {
loop {
let notified = self.watcher_idle.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.watcher_inflight.load(Ordering::SeqCst) == 0 {
return;
}
notified.await;
}
})
.await;
}
}
pub struct ReadyOnDrop {
states: Arc<IndexStates>,
project_id: u64,
}
impl Drop for ReadyOnDrop {
fn drop(&mut self) {
self.states.mark_ready(self.project_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_untouched_project_is_ready() {
let states = IndexStates::default();
assert_eq!(states.get(1), IndexState::Ready);
}
#[test]
fn building_reports_progress_and_percent() {
let states = IndexStates::default();
states.mark_building(1);
states.set_total(1, 10);
states.add_progress(1, 5);
let state = states.get(1);
assert!(state.is_building());
assert_eq!(state.percent(), Some(50));
states.mark_ready(1);
assert_eq!(states.get(1), IndexState::Ready);
}
#[test]
fn stale_tracks_pending_files() {
let states = IndexStates::default();
states.mark_stale(1, ["a.rs".to_string(), "b.rs".to_string()].into());
states.file_done(1, "a.rs");
let IndexState::Stale { pending } = states.get(1) else {
panic!("expected stale");
};
assert_eq!(pending.into_iter().collect::<Vec<_>>(), vec!["b.rs"]);
}
#[tokio::test]
async fn a_waiter_wakes_when_the_index_is_ready() {
let states = Arc::new(IndexStates::default());
states.mark_building(1);
let waiter = {
let states = Arc::clone(&states);
tokio::spawn(async move { states.wait_until_ready(1, Duration::from_secs(5)).await })
};
tokio::time::sleep(Duration::from_millis(20)).await;
drop(states.ready_on_drop(1));
assert_eq!(waiter.await.unwrap(), IndexState::Ready);
}
#[tokio::test]
async fn the_wait_is_bounded() {
let states = IndexStates::default();
states.mark_building(1);
let state = states.wait_until_ready(1, Duration::from_millis(30)).await;
assert!(state.is_building());
}
#[tokio::test]
async fn watcher_idle_waits_for_in_flight_batches() {
let states = Arc::new(IndexStates::default());
states.watcher_begin();
let ender = {
let states = Arc::clone(&states);
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(20)).await;
states.watcher_end();
})
};
states.wait_watcher_idle(Duration::from_secs(5)).await;
assert_eq!(states.watcher_inflight.load(Ordering::SeqCst), 0);
ender.await.unwrap();
}
}