use std::{borrow::Cow, collections::HashSet, hash::Hash, sync::Arc};
use async_trait::async_trait;
use pliantdb_core::schema::{view, CollectionName, Key, Schema, ViewName};
use pliantdb_jobs::{Job, Keyed};
use sled::{IVec, Transactional, Tree};
use super::{
mapper::{Map, Mapper},
view_document_map_tree_name, view_invalidated_docs_tree_name, view_versions_tree_name, Task,
};
use crate::database::{document_tree_name, Database};
#[derive(Debug)]
pub struct IntegrityScanner<DB> {
pub database: Database<DB>,
pub scan: IntegrityScan,
}
#[derive(Debug, Hash, Eq, PartialEq, Clone)]
pub struct IntegrityScan {
pub view_version: u64,
pub database: Arc<Cow<'static, str>>,
pub collection: CollectionName,
pub view_name: ViewName,
}
#[async_trait]
impl<DB> Job for IntegrityScanner<DB>
where
DB: Schema,
{
type Output = ();
async fn execute(&mut self) -> anyhow::Result<Self::Output> {
let documents = self.database.sled().open_tree(document_tree_name(
&self.database.data.name,
&self.scan.collection,
))?;
let view_versions = self.database.sled().open_tree(view_versions_tree_name(
&self.database.data.name,
&self.scan.collection,
))?;
let document_map = self.database.sled().open_tree(view_document_map_tree_name(
&self.database.data.name,
&self.scan.view_name,
))?;
let invalidated_entries =
self.database
.sled()
.open_tree(view_invalidated_docs_tree_name(
&self.database.data.name,
&self.scan.view_name,
))?;
let view_name = self.scan.view_name.clone();
let view_version = self.scan.view_version;
let needs_update = tokio::task::spawn_blocking::<_, anyhow::Result<bool>>(move || {
let document_ids = tree_keys::<u64>(&documents)?;
let view_is_current_version =
if let Some(version) = view_versions.get(view_name.to_string().as_bytes())? {
if let Ok(version) = u64::from_big_endian_bytes(&version) {
version == view_version
} else {
false
}
} else {
false
};
let missing_entries = if view_is_current_version {
let stored_document_ids = tree_keys::<u64>(&document_map)?;
document_ids
.difference(&stored_document_ids)
.cloned()
.collect::<HashSet<_>>()
} else {
document_ids
};
if !missing_entries.is_empty() {
(&invalidated_entries, &view_versions)
.transaction(|(invalidated_entries, view_versions)| {
view_versions.insert(
view_name.to_string().as_bytes(),
view_version.as_big_endian_bytes().unwrap().as_ref(),
)?;
for id in &missing_entries {
invalidated_entries.insert(
id.as_big_endian_bytes().unwrap().as_ref(),
IVec::default(),
)?;
}
Ok(())
})
.map_err(|err| match err {
sled::transaction::TransactionError::Abort(err) => err,
sled::transaction::TransactionError::Storage(err) => {
anyhow::Error::from(err)
}
})?;
return Ok(true);
}
Ok(false)
})
.await??;
if needs_update {
let job = self
.database
.data
.storage
.tasks()
.jobs
.lookup_or_enqueue(Mapper {
storage: self.database.clone(),
map: Map {
database: self.database.data.name.clone(),
collection: self.scan.collection.clone(),
view_name: self.scan.view_name.clone(),
},
})
.await;
job.receive().await?.map_err(crate::Error::Other)?;
}
self.database
.data
.storage
.tasks()
.mark_integrity_check_complete(
self.database.data.name.clone(),
self.scan.collection.clone(),
self.scan.view_name.clone(),
)
.await;
Ok(())
}
}
fn tree_keys<K: Key + Hash + Eq + Clone>(tree: &Tree) -> Result<HashSet<K>, anyhow::Error> {
let mut ids = HashSet::new();
for result in tree.iter() {
let (key, _) = result?;
let key = K::from_big_endian_bytes(&key).map_err(view::Error::KeySerialization)?;
ids.insert(key);
}
Ok(ids)
}
impl<DB> Keyed<Task> for IntegrityScanner<DB>
where
DB: Schema,
{
fn key(&self) -> Task {
Task::IntegrityScan(self.scan.clone())
}
}