use super::{Column, DataStore, PinModeRequirement};
use crate::error::Error;
use crate::repo::{PinKind, PinMode, PinStore, References};
use async_trait::async_trait;
use futures::stream::{StreamExt, TryStreamExt};
use libipld::cid::Cid;
use once_cell::sync::OnceCell;
use sled::{
self,
transaction::{
ConflictableTransactionError, TransactionError, TransactionResult, TransactionalTree,
UnabortableTransactionError,
},
Config as DbConfig, Db, Mode as DbMode,
};
use std::collections::BTreeSet;
use std::convert::Infallible;
use std::path::PathBuf;
use std::str::{self, FromStr};
#[derive(Debug)]
pub struct KvDataStore {
path: PathBuf,
db: OnceCell<Db>,
}
impl KvDataStore {
fn get_db(&self) -> &Db {
self.db.get().unwrap()
}
}
#[async_trait]
impl DataStore for KvDataStore {
fn new(root: PathBuf) -> KvDataStore {
KvDataStore {
path: root,
db: Default::default(),
}
}
async fn init(&self) -> Result<(), Error> {
let config = DbConfig::new();
let db = config
.mode(DbMode::HighThroughput)
.path(self.path.as_path())
.open()?;
match self.db.set(db) {
Ok(()) => Ok(()),
Err(_) => Err(anyhow::anyhow!("failed to init sled")),
}
}
async fn open(&self) -> Result<(), Error> {
Ok(())
}
async fn contains(&self, _col: Column, _key: &[u8]) -> Result<bool, Error> {
Err(anyhow::anyhow!("not implemented"))
}
async fn get(&self, _col: Column, _key: &[u8]) -> Result<Option<Vec<u8>>, Error> {
Err(anyhow::anyhow!("not implemented"))
}
async fn put(&self, _col: Column, _key: &[u8], _value: &[u8]) -> Result<(), Error> {
Err(anyhow::anyhow!("not implemented"))
}
async fn remove(&self, _col: Column, _key: &[u8]) -> Result<(), Error> {
Err(anyhow::anyhow!("not implemented"))
}
async fn wipe(&self) {
}
}
#[async_trait]
impl PinStore for KvDataStore {
async fn is_pinned(&self, cid: &Cid) -> Result<bool, Error> {
let cid = cid.to_owned();
let db = self.get_db().to_owned();
let span = tracing::Span::current();
tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
Ok(db.transaction::<_, _, Infallible>(|tree| {
Ok(get_pinned_mode(tree, &cid)?.is_some())
})?)
})
.await?
}
async fn insert_direct_pin(&self, target: &Cid) -> Result<(), Error> {
use ConflictableTransactionError::Abort;
let target = target.to_owned();
let db = self.get_db().to_owned();
let span = tracing::Span::current();
let res = tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
db.transaction(|tx_tree| {
let already_pinned = get_pinned_mode(tx_tree, &target)?;
match already_pinned {
Some((PinMode::Direct, _)) => return Ok(()),
Some((PinMode::Recursive, _)) => {
return Err(Abort(anyhow::anyhow!("already pinned recursively")))
}
Some((PinMode::Indirect, key)) => {
tx_tree.remove(key.as_str())?;
}
None => {}
}
let direct_key = get_pin_key(&target, &PinMode::Direct);
tx_tree.insert(direct_key.as_str(), direct_value())?;
tx_tree.flush();
Ok(())
})
})
.await?;
launder(res)
}
async fn insert_recursive_pin(
&self,
target: &Cid,
referenced: References<'_>,
) -> Result<(), Error> {
let set = referenced.try_collect::<BTreeSet<_>>().await?;
let target = target.to_owned();
let db = self.get_db().to_owned();
let span = tracing::Span::current();
tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
db.transaction::<_, _, Infallible>(move |tx_tree| {
let already_pinned = get_pinned_mode(tx_tree, &target)?;
match already_pinned {
Some((PinMode::Recursive, _)) => return Ok(()),
Some((PinMode::Direct, key)) | Some((PinMode::Indirect, key)) => {
tx_tree.remove(key.as_str())?;
}
None => {}
}
let recursive_key = get_pin_key(&target, &PinMode::Recursive);
tx_tree.insert(recursive_key.as_str(), recursive_value())?;
let target_value = indirect_value(&target);
for cid in set.iter() {
let indirect_key = get_pin_key(cid, &PinMode::Indirect);
if matches!(get_pinned_mode(tx_tree, cid)?, Some(_)) {
continue;
}
tx_tree.insert(indirect_key.as_str(), target_value.as_str())?;
}
tx_tree.flush();
Ok(())
})
})
.await??;
Ok(())
}
async fn remove_direct_pin(&self, target: &Cid) -> Result<(), Error> {
use ConflictableTransactionError::Abort;
let target = target.to_owned();
let db = self.get_db().to_owned();
let span = tracing::Span::current();
let res = tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
db.transaction::<_, _, Error>(|tx_tree| {
if is_not_pinned_or_pinned_indirectly(tx_tree, &target)? {
return Err(Abort(anyhow::anyhow!("not pinned or pinned indirectly")));
}
let key = get_pin_key(&target, &PinMode::Direct);
tx_tree.remove(key.as_str())?;
tx_tree.flush();
Ok(())
})
})
.await?;
launder(res)
}
async fn remove_recursive_pin(
&self,
target: &Cid,
referenced: References<'_>,
) -> Result<(), Error> {
use ConflictableTransactionError::Abort;
let set = referenced.try_collect::<BTreeSet<_>>().await?;
let target = target.to_owned();
let db = self.get_db().to_owned();
let span = tracing::Span::current();
let res = tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
db.transaction(|tx_tree| {
if is_not_pinned_or_pinned_indirectly(tx_tree, &target)? {
return Err(Abort(anyhow::anyhow!("not pinned or pinned indirectly")));
}
let recursive_key = get_pin_key(&target, &PinMode::Recursive);
tx_tree.remove(recursive_key.as_str())?;
for cid in &set {
let already_pinned = get_pinned_mode(tx_tree, cid)?;
match already_pinned {
Some((PinMode::Recursive, _)) | Some((PinMode::Direct, _)) => continue, Some((PinMode::Indirect, key)) => {
tx_tree.remove(key.as_str())?;
}
None => {}
}
}
tx_tree.flush();
Ok(())
})
})
.await?;
launder(res)
}
async fn list(
&self,
requirement: Option<PinMode>,
) -> futures::stream::BoxStream<'static, Result<(Cid, PinMode), Error>> {
use tokio_stream::wrappers::UnboundedReceiverStream;
let db = self.get_db().to_owned();
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
let span = tracing::Span::current();
let _jh = tokio::task::spawn_blocking(move || {
let span = tracing::trace_span!(parent: &span, "blocking");
let _g = span.enter();
let iter = db.range::<String, std::ops::RangeFull>(..);
let requirement = PinModeRequirement::from(requirement);
let adapted =
iter.map(|res| res.map_err(Error::from))
.filter_map(move |res| match res {
Ok((k, _v)) => {
if !k.starts_with(b"pin.") || k.len() < 7 {
return Some(Err(anyhow::anyhow!(
"invalid pin: {:?}",
&*String::from_utf8_lossy(&k)
)));
}
let mode = match k[4] {
b'd' => PinMode::Direct,
b'r' => PinMode::Recursive,
b'i' => PinMode::Indirect,
x => {
return Some(Err(anyhow::anyhow!(
"invalid pinmode: {}",
x as char
)))
}
};
if !requirement.matches(&mode) {
None
} else {
let cid = std::str::from_utf8(&k[6..]).map_err(Error::from);
let cid = cid.and_then(|x| Cid::from_str(x).map_err(Error::from));
let cid = cid.map_err(|e| {
e.context(format!(
"failed to read pin: {:?}",
&*String::from_utf8_lossy(&k)
))
});
Some(cid.map(move |cid| (cid, mode)))
}
}
Err(e) => Some(Err(e)),
});
for res in adapted {
if tx.send(res).is_err() {
break;
}
}
});
UnboundedReceiverStream::new(rx).boxed()
}
async fn query(
&self,
ids: Vec<Cid>,
requirement: Option<PinMode>,
) -> Result<Vec<(Cid, PinKind<Cid>)>, Error> {
use ConflictableTransactionError::Abort;
let requirement = PinModeRequirement::from(requirement);
let db = self.get_db().to_owned();
tokio::task::spawn_blocking(move || {
let res = db.transaction::<_, _, Error>(|tx_tree| {
let mut modes = Vec::with_capacity(ids.len());
for id in ids.iter() {
let mode_and_key = get_pinned_mode(tx_tree, id)?;
let matched = match mode_and_key {
Some((pin_mode, key)) if requirement.matches(&pin_mode) => match pin_mode {
PinMode::Direct => Some(PinKind::Direct),
PinMode::Recursive => Some(PinKind::Recursive(0)),
PinMode::Indirect => tx_tree
.get(key.as_str())?
.map(|root| {
cid_from_indirect_value(&root)
.map(PinKind::IndirectFrom)
.map_err(|e| {
Abort(e.context(format!(
"failed to read indirect pin source: {:?}",
String::from_utf8_lossy(root.as_ref()).as_ref(),
)))
})
})
.transpose()?,
},
Some(_) | None => None,
};
modes.push(matched);
}
Ok(modes)
});
let modes = launder(res)?;
Ok(ids
.into_iter()
.zip(modes.into_iter())
.filter_map(|(cid, mode)| mode.map(move |mode| (cid, mode)))
.collect::<Vec<_>>())
})
.await?
}
}
fn direct_value() -> &'static [u8] {
Default::default()
}
fn recursive_value() -> &'static [u8] {
Default::default()
}
fn indirect_value(recursively_pinned: &Cid) -> String {
recursively_pinned.to_string()
}
fn cid_from_indirect_value(bytes: &[u8]) -> Result<Cid, Error> {
str::from_utf8(bytes)
.map_err(Error::from)
.and_then(|s| Cid::from_str(s).map_err(Error::from))
}
fn launder<T>(res: TransactionResult<T, Error>) -> Result<T, Error> {
use TransactionError::*;
match res {
Ok(t) => Ok(t),
Err(Abort(e)) => Err(e),
Err(Storage(e)) => Err(e.into()),
}
}
fn pin_mode_literal(pin_mode: &PinMode) -> &'static str {
match pin_mode {
PinMode::Direct => "d",
PinMode::Indirect => "i",
PinMode::Recursive => "r",
}
}
fn get_pin_key(cid: &Cid, pin_mode: &PinMode) -> String {
format!("pin.{}.{}", pin_mode_literal(pin_mode), cid)
}
fn get_pinned_mode(
tree: &TransactionalTree,
block: &Cid,
) -> Result<Option<(PinMode, String)>, UnabortableTransactionError> {
for mode in &[PinMode::Direct, PinMode::Recursive, PinMode::Indirect] {
let key = get_pin_key(block, mode);
if tree.get(key.as_str())?.is_some() {
return Ok(Some((*mode, key)));
}
}
Ok(None)
}
fn is_not_pinned_or_pinned_indirectly(
tree: &TransactionalTree,
block: &Cid,
) -> Result<bool, UnabortableTransactionError> {
match get_pinned_mode(tree, block)? {
Some((PinMode::Indirect, _)) | None => Ok(true),
_ => Ok(false),
}
}
#[cfg(test)]
crate::pinstore_interface_tests!(common_tests, crate::repo::kv::KvDataStore::new);