#![cfg(feature = "kv-tikv")]
mod cnf;
use crate::err::Error;
use crate::key::debug::Sprintable;
use crate::kvs::savepoint::{SaveOperation, SavePointImpl, SavePoints, SavePrepare};
use crate::kvs::Check;
use crate::kvs::Key;
use crate::kvs::KeyEncode;
use crate::kvs::Val;
use crate::vs::VersionStamp;
use std::fmt::Debug;
use std::ops::Range;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use tikv::TimestampExt;
use tikv::TransactionOptions;
use tikv::{CheckLevel, Config, TransactionClient};
const TARGET: &str = "surrealdb::core::kvs::tikv";
pub struct Datastore {
db: Pin<Arc<TransactionClient>>,
}
pub struct Transaction {
done: bool,
write: bool,
check: Check,
inner: tikv::Transaction,
save_points: SavePoints,
db: Pin<Arc<TransactionClient>>,
}
impl Drop for Transaction {
fn drop(&mut self) {
if !self.done && self.write {
match self.check {
Check::None => {
trace!("A transaction was dropped without being committed or cancelled");
}
Check::Warn => {
warn!("A transaction was dropped without being committed or cancelled");
}
Check::Error => {
error!("A transaction was dropped without being committed or cancelled");
}
}
}
}
}
impl Datastore {
pub(crate) async fn new(path: &str) -> Result<Datastore, Error> {
let config = match *cnf::TIKV_API_VERSION {
2 => match *cnf::TIKV_KEYSPACE {
Some(ref keyspace) => {
info!(target: TARGET, "Connecting to keyspace with cluster API V2: {keyspace}");
Config::default().with_keyspace(keyspace)
}
None => {
info!(target: TARGET, "Connecting to default keyspace with cluster API V2");
Config::default().with_default_keyspace()
}
},
1 => {
info!(target: TARGET, "Connecting with cluster API V1");
Config::default()
}
_ => return Err(Error::Ds("Invalid TiKV API version".into())),
};
let config = config.with_timeout(Duration::from_secs(*cnf::TIKV_REQUEST_TIMEOUT));
let config =
config.with_grpc_max_decoding_message_size(*cnf::TIKV_GRPC_MAX_DECODING_MESSAGE_SIZE);
let client = TransactionClient::new_with_config(vec![path], config);
match client.await {
Ok(db) => Ok(Datastore {
db: Arc::pin(db),
}),
Err(e) => Err(Error::Ds(e.to_string())),
}
}
pub(crate) async fn shutdown(&self) -> Result<(), Error> {
Ok(())
}
pub(crate) async fn transaction(&self, write: bool, lock: bool) -> Result<Transaction, Error> {
let mut opt = if lock {
TransactionOptions::new_pessimistic()
} else {
TransactionOptions::new_optimistic()
};
opt = match *cnf::TIKV_ASYNC_COMMIT {
true => opt.use_async_commit(),
_ => opt,
};
opt = match *cnf::TIKV_ONE_PHASE_COMMIT {
true => opt.try_one_pc(),
_ => opt,
};
opt = opt.drop_check(CheckLevel::Warn);
if !write {
opt = opt.read_only();
}
#[cfg(not(debug_assertions))]
let check = Check::Warn;
#[cfg(debug_assertions)]
let check = Check::Error;
match self.db.begin_with_options(opt).await {
Ok(inner) => Ok(Transaction {
done: false,
check,
write,
inner,
db: self.db.clone(),
save_points: Default::default(),
}),
Err(e) => Err(Error::Tx(e.to_string())),
}
}
}
impl super::api::Transaction for Transaction {
fn supports_reverse_scan(&self) -> bool {
true
}
fn check_level(&mut self, check: Check) {
self.check = check;
}
fn closed(&self) -> bool {
self.done
}
fn writeable(&self) -> bool {
self.write
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self))]
async fn cancel(&mut self) -> Result<(), Error> {
if self.done {
return Err(Error::TxFinished);
}
self.done = true;
if self.write {
self.inner.rollback().await?;
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self))]
async fn commit(&mut self) -> Result<(), Error> {
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
self.done = true;
if let Err(err) = self.inner.commit().await {
if let Err(inner_err) = self.inner.rollback().await {
error!("Transaction commit failed {} and rollback failed: {}", err, inner_err);
}
return Err(err.into());
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn exists<K>(&mut self, key: K, version: Option<u64>) -> Result<bool, Error>
where
K: KeyEncode + Sprintable + Debug,
{
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.done {
return Err(Error::TxFinished);
}
let res = self.inner.key_exists(key.encode_owned()?).await?;
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn get<K>(&mut self, key: K, version: Option<u64>) -> Result<Option<Val>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.done {
return Err(Error::TxFinished);
}
let res = self.inner.get(key.encode_owned()?).await?;
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn set<K, V>(&mut self, key: K, val: V, version: Option<u64>) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
V: Into<Val> + Debug,
{
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
let key = key.encode_owned()?;
let prep = if self.save_points.is_some() {
self.save_point_prepare(&key, version, SaveOperation::Set).await?
} else {
None
};
self.inner.put(key, val.into()).await?;
if let Some(prep) = prep {
self.save_points.save(prep);
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn put<K, V>(&mut self, key: K, val: V, version: Option<u64>) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
V: Into<Val> + Debug,
{
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
let key = key.encode_owned()?;
let val = val.into();
let prep = if self.save_points.is_some() {
self.save_point_prepare(&key, version, SaveOperation::Put).await?
} else {
None
};
let key_exists = if let Some(SavePrepare::NewKey(_, sv)) = &prep {
sv.get_val().is_some()
} else {
self.inner.key_exists(key.clone()).await?
};
if key_exists {
return Err(Error::TxKeyAlreadyExists);
}
self.inner.put(key, val).await?;
if let Some(prep) = prep {
self.save_points.save(prep);
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn putc<K, V>(&mut self, key: K, val: V, chk: Option<V>) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
V: Into<Val> + Debug,
{
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
let key = key.encode_owned()?;
let val = val.into();
let chk = chk.map(Into::into);
let prep = if self.save_points.is_some() {
self.save_point_prepare(&key, None, SaveOperation::Put).await?
} else {
None
};
let current_val = if let Some(SavePrepare::NewKey(_, sv)) = &prep {
sv.get_val().cloned()
} else {
self.inner.get(key.clone()).await?
};
match (current_val, chk) {
(Some(v), Some(w)) if v == w => self.inner.put(key, val).await?,
(None, None) => self.inner.put(key, val).await?,
_ => return Err(Error::TxConditionNotMet),
};
if let Some(prep) = prep {
self.save_points.save(prep);
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn del<K>(&mut self, key: K) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
{
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
let key = key.encode_owned()?;
let prep = if self.save_points.is_some() {
self.save_point_prepare(&key, None, SaveOperation::Del).await?
} else {
None
};
self.inner.delete(key).await?;
if let Some(prep) = prep {
self.save_points.save(prep);
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn delc<K, V>(&mut self, key: K, chk: Option<V>) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
V: Into<Val> + Debug,
{
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
let key = key.encode_owned()?;
let chk = chk.map(Into::into);
let prep = if self.save_points.is_some() {
self.save_point_prepare(&key, None, SaveOperation::Del).await?
} else {
None
};
let current_val = if let Some(SavePrepare::NewKey(_, sv)) = &prep {
sv.get_val().cloned()
} else {
self.inner.get(key.clone()).await?
};
match (current_val, chk) {
(Some(v), Some(w)) if v == w => self.inner.delete(key).await?,
(None, None) => self.inner.delete(key).await?,
_ => return Err(Error::TxConditionNotMet),
};
if let Some(prep) = prep {
self.save_points.save(prep);
}
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
async fn delr<K>(&mut self, rng: Range<K>) -> Result<(), Error>
where
K: KeyEncode + Sprintable + Debug,
{
if self.done {
return Err(Error::TxFinished);
}
if !self.write {
return Err(Error::TxReadonly);
}
self.db.unsafe_destroy_range(rng.start.encode_owned()?..rng.end.encode_owned()?).await?;
Ok(())
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
async fn keys<K>(
&mut self,
rng: Range<K>,
limit: u32,
version: Option<u64>,
) -> Result<Vec<Key>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
let rng = self.prepare_scan(rng, version)?;
let res = self.inner.scan_keys(rng, limit).await?.map(Key::from).collect();
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
async fn keysr<K>(
&mut self,
rng: Range<K>,
limit: u32,
version: Option<u64>,
) -> Result<Vec<Key>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
let rng = self.prepare_scan(rng, version)?;
let res = self.inner.scan_keys_reverse(rng, limit).await?.map(Key::from).collect();
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
async fn scan<K>(
&mut self,
rng: Range<K>,
limit: u32,
version: Option<u64>,
) -> Result<Vec<(Key, Val)>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
let rng = self.prepare_scan(rng, version)?;
let res = self.inner.scan(rng, limit).await?.map(|kv| (Key::from(kv.0), kv.1)).collect();
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
async fn scanr<K>(
&mut self,
rng: Range<K>,
limit: u32,
version: Option<u64>,
) -> Result<Vec<(Key, Val)>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
let rng = self.prepare_scan(rng, version)?;
let res =
self.inner.scan_reverse(rng, limit).await?.map(|kv| (Key::from(kv.0), kv.1)).collect();
Ok(res)
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
async fn get_timestamp<K>(&mut self, key: K) -> Result<VersionStamp, Error>
where
K: KeyEncode + Sprintable + Debug,
{
if self.done {
return Err(Error::TxFinished);
}
let key = key.encode_owned()?;
let ver = self.inner.current_timestamp().await?.version();
if let Some(prev) = self.get(key.as_slice(), None).await? {
let prev = VersionStamp::from_slice(prev.as_slice())?.try_into_u64()?;
if prev >= ver {
return Err(Error::TxFailure);
}
};
let ver = VersionStamp::from_u64(ver);
self.set(key.as_slice(), ver.as_bytes(), None).await?;
Ok(ver)
}
}
impl SavePointImpl for Transaction {
fn get_save_points(&mut self) -> &mut SavePoints {
&mut self.save_points
}
}
impl Transaction {
fn prepare_scan<K>(&self, rng: Range<K>, version: Option<u64>) -> Result<Range<Key>, Error>
where
K: KeyEncode + Sprintable + Debug,
{
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.done {
return Err(Error::TxFinished);
}
let rng: Range<Key> = Range {
start: rng.start.encode_owned()?,
end: rng.end.encode_owned()?,
};
Ok(rng)
}
}