#![cfg(feature = "kv-tikv")]
mod cnf;
mod savepoint;
use std::collections::HashMap;
use std::ops::Range;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use chrono::{DateTime, Utc};
pub use cnf::TikvConfig;
use savepoint::{Operation, Savepoint};
use tikv::transaction::ResolveLocksOptions;
use tikv::{CheckLevel, Config, TimestampExt, TransactionClient, TransactionOptions};
use tokio::sync::RwLock;
use super::api::{BoxFut, GetMultiResult, KeysResult, ScanResult};
use super::err::{Error, Result};
use super::timestamp::MAX_TIMESTAMP_BYTES;
use super::util;
use crate::key::debug::Sprintable;
use crate::kvs::api::Transactable;
use crate::kvs::timestamp::{BoxTimeStamp, BoxTimeStampImpl};
use crate::kvs::{COUNT_BATCH_SIZE, Key, TimeStamp, TimeStampImpl, Val};
const TARGET: &str = "surrealdb::core::kvs::tikv";
pub struct TikvOpsHandle {
db: Pin<Arc<TransactionClient>>,
in_flight_txns: Arc<AtomicUsize>,
config: TikvConfig,
}
pub struct Datastore {
handle: Arc<TikvOpsHandle>,
}
pub struct Transaction {
done: AtomicBool,
write: bool,
inner: RwLock<TransactionInner>,
started_at: Instant,
#[allow(dead_code, reason = "Held to keep the TransactionClient alive while `tx` borrows it")]
handle: Arc<TikvOpsHandle>,
}
impl Transaction {
fn release_in_flight(&self) {
let _ = self
.handle
.in_flight_txns
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |n| n.checked_sub(1));
}
async fn delete_range_bounded(&self, rng: Range<Key>) -> Result<()> {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let max_keys = self.handle.config.delr_max_keys;
let mut inner = self.inner.write().await;
let track_ops = !inner.savepoints.is_empty() || !inner.operations.is_empty();
let end = rng.end.clone();
let mut start = rng.start;
let mut processed: u32 = 0;
let mut previous_batch_full = false;
loop {
let remaining = max_keys.saturating_sub(processed);
if remaining == 0 {
if previous_batch_full {
let mut iter = inner.tx.scan_keys(start.clone()..end.clone(), 1).await?;
if iter.next().is_some() {
return Err(Error::TransactionRangeTooLarge(max_keys));
}
}
break;
}
let batch_size = remaining.min(cnf::TIKV_DELR_BATCH_SIZE);
let keys: Vec<tikv::Key> =
inner.tx.scan_keys(start.clone()..end.clone(), batch_size).await?.collect();
if keys.is_empty() {
break;
}
previous_batch_full = (keys.len() as u32) == batch_size;
let last = keys.last().cloned();
for k in keys {
let key = Key::from(k);
let old_val = if track_ops {
inner.tx.get(key.clone()).await?
} else {
None
};
inner.tx.delete(key.clone()).await?;
if let Some(val) = old_val {
inner.operations.push(Operation::RestoreDeleted(key, val));
}
processed = processed.saturating_add(1);
}
match last {
Some(k) => {
let mut next = Key::from(k);
util::advance_key(&mut next);
start = next;
}
None => break,
}
}
Ok(())
}
}
impl Drop for Transaction {
fn drop(&mut self) {
if !self.done.swap(true, Ordering::AcqRel) {
self.release_in_flight();
}
}
}
struct TransactionInner {
tx: tikv::Transaction,
savepoints: Vec<Savepoint>,
operations: Vec<Operation>,
}
impl Datastore {
pub(crate) async fn new(path: &str, tikv_config: TikvConfig) -> Result<Datastore> {
let config = match tikv_config.api_version {
2 => match tikv_config.keyspace.as_ref() {
Some(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::Datastore("Invalid TiKV API version".into())),
};
let config = config.with_timeout(Duration::from_secs(tikv_config.request_timeout_secs));
let config = config
.with_grpc_max_decoding_message_size(tikv_config.grpc_max_decoding_message_size)
.with_grpc_max_encoding_message_size(tikv_config.grpc_max_encoding_message_size);
let config = match validate_tls_paths(&tikv_config)? {
Some(TlsPaths {
ca,
cert,
key,
}) => {
info!(
target: TARGET,
ca = %ca,
cert = %cert,
key = %key,
"Configuring TiKV client with mTLS",
);
config.with_security(ca, cert, key)
}
None => config,
};
let db = match TransactionClient::new_with_config(vec![path], config).await {
Ok(db) => Arc::pin(db),
Err(e) => return Err(Error::Datastore(e.to_string())),
};
if tikv_config.health_probe {
let probe = async {
db.current_timestamp().await.map_err(Error::from)?;
let mut txn = db
.begin_with_options(TransactionOptions::new_optimistic().read_only())
.await
.map_err(Error::from)?;
let mut iter = txn.scan_keys(vec![0u8]..vec![0u8], 1).await.map_err(Error::from)?;
while iter.next().is_some() {}
Ok::<_, Error>(())
};
let probe_budget =
Duration::from_secs(tikv_config.request_timeout_secs.saturating_mul(2));
match tokio::time::timeout(probe_budget, probe).await {
Ok(Ok(())) => {
info!(target: TARGET, "TiKV startup health probe succeeded");
}
Ok(Err(e)) => {
return Err(Error::Datastore(format!("TiKV startup health probe failed: {e}")));
}
Err(_) => {
return Err(Error::Datastore(format!(
"TiKV startup health probe timed out after {probe_budget:?}"
)));
}
}
}
let handle = Arc::new(TikvOpsHandle {
db,
in_flight_txns: Arc::new(AtomicUsize::new(0)),
config: tikv_config,
});
Ok(Datastore {
handle,
})
}
pub(crate) fn ops_handle(&self) -> Arc<TikvOpsHandle> {
Arc::clone(&self.handle)
}
pub(crate) async fn shutdown(&self) -> Result<()> {
let cfg = &self.handle.config;
let drain_deadline = Instant::now() + Duration::from_secs(cfg.shutdown_grace_secs);
let drained = loop {
let outstanding = self.handle.in_flight_transaction_count();
if outstanding == 0 {
break true;
}
if Instant::now() >= drain_deadline {
warn!(
target: TARGET,
outstanding,
"TiKV shutdown drain timed out; proceeding with active transactions still in flight",
);
break false;
}
tokio::time::sleep(Duration::from_millis(100)).await;
};
if drained && cfg.gc_enabled {
let shutdown_gc_timeout = Duration::from_secs(cfg.shutdown_gc_timeout_secs);
let gc_lifetime = Duration::from_secs(cfg.gc_lifetime_secs);
match tokio::time::timeout(shutdown_gc_timeout, self.handle.run_mvcc_gc(gc_lifetime))
.await
{
Ok(Ok(())) => {}
Ok(Err(e)) => {
warn!(
target: TARGET,
error = %e,
"Advisory TiKV GC pass at shutdown failed",
);
}
Err(_) => {
warn!(
target: TARGET,
timeout_ms = shutdown_gc_timeout.as_millis() as u64,
"Advisory TiKV GC pass at shutdown timed out",
);
}
}
}
Ok(())
}
pub(crate) async fn transaction(
&self,
write: bool,
lock: bool,
) -> Result<Box<dyn Transactable>> {
let cfg = &self.handle.config;
let mut opt = if lock {
TransactionOptions::new_pessimistic()
} else {
TransactionOptions::new_optimistic()
};
if cfg.async_commit {
opt = opt.use_async_commit();
}
if cfg.one_phase_commit {
opt = opt.try_one_pc();
}
opt = opt.drop_check(if cfg!(debug_assertions) {
CheckLevel::Panic
} else {
CheckLevel::Warn
});
if !write {
opt = opt.read_only();
}
match self.handle.db.begin_with_options(opt).await {
Ok(txn) => {
self.handle.in_flight_txns.fetch_add(1, Ordering::AcqRel);
Ok(Box::new(Transaction {
done: AtomicBool::new(false),
write,
inner: RwLock::new(TransactionInner {
tx: txn,
savepoints: Vec::new(),
operations: Vec::new(),
}),
started_at: Instant::now(),
handle: Arc::clone(&self.handle),
}))
}
Err(e) => Err(Error::from(e)),
}
}
}
impl TikvOpsHandle {
pub fn in_flight_transaction_count(&self) -> usize {
self.in_flight_txns.load(Ordering::Acquire)
}
pub async fn unsafe_destroy_range(&self, start: Vec<u8>, end: Vec<u8>) -> Result<()> {
let started = Instant::now();
let tikv_start: tikv::Key = start.into();
let tikv_end: tikv::Key = end.into();
self.db.unsafe_destroy_range(tikv_start..tikv_end).await.map_err(Error::from)?;
debug!(
target: TARGET,
duration_ms = started.elapsed().as_millis() as u64,
"TiKV unsafe_destroy_range completed",
);
Ok(())
}
pub async fn run_mvcc_gc(&self, lifetime: Duration) -> Result<()> {
if !self.config.gc_enabled {
return Ok(());
}
let started = Instant::now();
let now = self.db.current_timestamp().await.map_err(Error::from)?;
let safepoint = match safepoint_from(&now, lifetime) {
Some(ts) => ts,
None => {
debug!(
target: TARGET,
"Skipping TiKV MVCC GC pass: safepoint would precede epoch",
);
return Ok(());
}
};
let advanced = match self.db.gc(safepoint).await {
Ok(v) => v,
Err(e) => {
warn!(
target: TARGET,
error = %e,
"TiKV MVCC GC pass failed",
);
return Err(Error::from(e));
}
};
info!(
target: TARGET,
advanced,
duration_ms = started.elapsed().as_millis() as u64,
lifetime_secs = lifetime.as_secs(),
"TiKV MVCC GC pass completed",
);
Ok(())
}
pub async fn run_lock_cleanup(&self, lifetime: Duration) -> Result<()> {
if !self.config.gc_enabled {
return Ok(());
}
let started = Instant::now();
let now = self.db.current_timestamp().await.map_err(Error::from)?;
let safepoint = match safepoint_from(&now, lifetime) {
Some(ts) => ts,
None => {
debug!(
target: TARGET,
"Skipping TiKV lock cleanup pass: safepoint would precede epoch",
);
return Ok(());
}
};
let range: Range<tikv::Key> = (vec![0u8].into())..(vec![0xffu8; 16].into());
let result = self
.db
.cleanup_locks(range, &safepoint, ResolveLocksOptions::default())
.await
.map_err(Error::from)?;
info!(
target: TARGET,
has_region_error = result.region_error.is_some(),
key_error_count = result.key_error.as_ref().map(|v| v.len()).unwrap_or(0),
resolved_locks = result.resolved_locks,
duration_ms = started.elapsed().as_millis() as u64,
lifetime_secs = lifetime.as_secs(),
"TiKV lock cleanup pass completed",
);
Ok(())
}
}
fn safepoint_from(now: &tikv::Timestamp, lifetime: Duration) -> Option<tikv::Timestamp> {
let millis: i64 = lifetime.as_millis().try_into().ok()?;
let physical = now.physical.checked_sub(millis)?;
if physical < 0 {
return None;
}
Some(tikv::Timestamp {
physical,
logical: now.logical,
suffix_bits: now.suffix_bits,
})
}
#[cfg_attr(test, derive(Debug))]
struct TlsPaths {
ca: String,
cert: String,
key: String,
}
fn validate_tls_paths(config: &TikvConfig) -> Result<Option<TlsPaths>> {
match (config.tls_ca_path.as_ref(), config.tls_cert_path.as_ref(), config.tls_key_path.as_ref())
{
(Some(ca), Some(cert), Some(key)) => Ok(Some(TlsPaths {
ca: ca.clone(),
cert: cert.clone(),
key: key.clone(),
})),
(None, None, None) => Ok(None),
(ca, cert, key) => Err(Error::Datastore(format!(
"TiKV mTLS requires tikv_tls_ca_path, tikv_tls_cert_path and \
tikv_tls_key_path to all be set together (e.g. via \
SURREAL_TIKV_TLS_*); got ca={}, cert={}, key={}",
if ca.is_some() {
"set"
} else {
"unset"
},
if cert.is_some() {
"set"
} else {
"unset"
},
if key.is_some() {
"set"
} else {
"unset"
},
))),
}
}
impl Transactable for Transaction {
fn kind(&self) -> &'static str {
"tikv"
}
fn closed(&self) -> bool {
self.done.load(Ordering::Relaxed)
}
fn writeable(&self) -> bool {
self.write
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self))]
fn cancel(&self) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.done.swap(true, Ordering::AcqRel) {
return Err(Error::TransactionFinished);
}
self.release_in_flight();
let mut inner = self.inner.write().await;
if self.write {
let _ = inner.tx.rollback().await;
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self))]
fn commit(&self) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.done.swap(true, Ordering::AcqRel) {
return Err(Error::TransactionFinished);
}
self.release_in_flight();
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
if let Err(err) = inner.tx.commit().await {
if let Err(inner_err) = inner.tx.rollback().await {
error!(
target: TARGET,
commit_error = %err,
rollback_error = %inner_err,
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
"TiKV transaction commit failed and rollback also failed",
);
} else {
debug!(
target: TARGET,
commit_error = %err,
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
"TiKV transaction commit failed; rollback succeeded",
);
}
return Err(err.into());
}
trace!(
target: TARGET,
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
"TiKV transaction committed",
);
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn exists(&self, key: Key, version: Option<u64>) -> BoxFut<'_, Result<bool>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let res = inner.tx.key_exists(key).await?;
Ok(res)
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn get(&self, key: Key, version: Option<u64>) -> BoxFut<'_, Result<Option<Val>>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let res = inner.tx.get(key).await?;
Ok(res)
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(keys = keys.sprint()))]
fn getm(&self, keys: Vec<Key>, version: Option<u64>) -> BoxFut<'_, Result<GetMultiResult>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let mut key_index: HashMap<&[u8], Vec<usize>> = HashMap::with_capacity(keys.len());
for (i, k) in keys.iter().enumerate() {
key_index.entry(k.as_slice()).or_default().push(i);
}
let pairs = inner.tx.batch_get(keys.iter().cloned()).await?;
let mut values: Vec<Option<Val>> = vec![None; keys.len()];
let mut records = 0u64;
let mut value_bytes = 0u64;
for kv in pairs {
if let Some(idxs) = key_index.get(Key::from(kv.0).as_slice())
&& let Some((&last, rest)) = idxs.split_last()
{
let len = kv.1.len() as u64;
for &i in rest {
records += 1;
value_bytes += len;
values[i] = Some(kv.1.clone());
}
records += 1;
value_bytes += len;
values[last] = Some(kv.1);
}
}
Ok(GetMultiResult {
values,
records,
value_bytes,
})
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn set(&self, key: Key, val: Val) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let old_val = if !inner.savepoints.is_empty() || !inner.operations.is_empty() {
inner.tx.get(key.clone()).await?
} else {
None
};
inner.tx.put(key.clone(), val).await?;
if !inner.savepoints.is_empty() || !inner.operations.is_empty() {
match old_val {
Some(existing_val) => {
inner.operations.push(Operation::RestoreValue(key, existing_val));
}
None => {
inner.operations.push(Operation::DeleteKey(key));
}
}
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn put(&self, key: Key, val: Val) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let exists = inner.tx.key_exists(key.clone()).await?;
if exists {
return Err(Error::TransactionKeyAlreadyExists);
}
inner.tx.put(key.clone(), val).await?;
if !inner.savepoints.is_empty() || !inner.operations.is_empty() {
inner.operations.push(Operation::DeleteKey(key));
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn putc(&self, key: Key, val: Val, chk: Option<Val>) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let current = inner.tx.get(key.clone()).await?;
match (¤t, &chk) {
(Some(v), Some(w)) if v == w => {}
(None, None) => {}
_ => return Err(Error::TransactionConditionNotMet),
};
inner.tx.put(key.clone(), val).await?;
if !inner.savepoints.is_empty() || !inner.operations.is_empty() {
match current {
Some(existing_val) => {
inner.operations.push(Operation::RestoreValue(key, existing_val));
}
None => {
inner.operations.push(Operation::DeleteKey(key));
}
}
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn del(&self, key: Key) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let old_val = if !inner.savepoints.is_empty() || !inner.operations.is_empty() {
inner.tx.get(key.clone()).await?
} else {
None
};
inner.tx.delete(key.clone()).await?;
if let Some(existing_val) = old_val {
inner.operations.push(Operation::RestoreDeleted(key, existing_val));
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(key = key.sprint()))]
fn delc(&self, key: Key, chk: Option<Val>) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let current = inner.tx.get(key.clone()).await?;
match (¤t, &chk) {
(Some(v), Some(w)) if v == w => {}
(None, None) => {}
_ => return Err(Error::TransactionConditionNotMet),
};
inner.tx.delete(key.clone()).await?;
if let Some(existing_val) = current {
inner.operations.push(Operation::RestoreDeleted(key, existing_val));
}
Ok(())
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn delr(&self, rng: Range<Key>) -> BoxFut<'_, Result<()>> {
Box::pin(async move { self.delete_range_bounded(rng).await })
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn clrr(&self, rng: Range<Key>) -> BoxFut<'_, Result<()>> {
Box::pin(async move { self.delete_range_bounded(rng).await })
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn count(&self, rng: Range<Key>, version: Option<u64>) -> BoxFut<'_, Result<usize>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let mut total = 0usize;
let end = rng.end.clone();
let mut start = rng.start;
loop {
let iter = inner.tx.scan_keys(start..end.clone(), COUNT_BATCH_SIZE).await?;
let mut key: Option<tikv::Key> = None;
let mut count = 0u32;
for k in iter {
count += 1;
key = Some(k);
}
total += count as usize;
if count < COUNT_BATCH_SIZE {
break;
}
match key {
Some(k) => {
let mut k = Key::from(k);
util::advance_key(&mut k);
start = k;
}
None => break,
}
}
Ok(total)
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn keys(
&self,
rng: Range<Key>,
limit: u32,
skip: u32,
version: Option<u64>,
) -> BoxFut<'_, Result<KeysResult>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let count = limit.saturating_add(skip);
let mut iter = inner.tx.scan_keys(rng, count).await?;
Ok(consume_keys(&mut iter, limit, skip))
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn keysr(
&self,
rng: Range<Key>,
limit: u32,
skip: u32,
version: Option<u64>,
) -> BoxFut<'_, Result<KeysResult>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let count = limit.saturating_add(skip);
let mut iter = inner.tx.scan_keys_reverse(rng, count).await?;
Ok(consume_keys(&mut iter, limit, skip))
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn scan(
&self,
rng: Range<Key>,
limit: u32,
skip: u32,
version: Option<u64>,
) -> BoxFut<'_, Result<ScanResult>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let rng = if skip > 0 {
let skipped = inner.tx.scan_keys(rng.clone(), skip).await?;
match skipped.last() {
Some(last) => {
let mut start: Key = Key::from(last);
util::advance_key(&mut start);
start..rng.end
}
None => return Ok(ScanResult::default()),
}
} else {
rng
};
let mut iter = inner.tx.scan(rng, limit).await?;
Ok(consume_vals(&mut iter, limit))
})
}
#[instrument(level = "trace", target = "surrealdb::core::kvs::api", skip(self), fields(rng = rng.sprint()))]
fn scanr(
&self,
rng: Range<Key>,
limit: u32,
skip: u32,
version: Option<u64>,
) -> BoxFut<'_, Result<ScanResult>> {
Box::pin(async move {
if version.is_some() {
return Err(Error::UnsupportedVersionedQueries);
}
if self.closed() {
return Err(Error::TransactionFinished);
}
let mut inner = self.inner.write().await;
let rng = if skip > 0 {
let skipped = inner.tx.scan_keys_reverse(rng.clone(), skip).await?;
match skipped.last() {
Some(last) => {
let end: Key = Key::from(last);
rng.start..end
}
None => return Ok(ScanResult::default()),
}
} else {
rng
};
let mut iter = inner.tx.scan_reverse(rng, limit).await?;
Ok(consume_vals(&mut iter, limit))
})
}
fn new_save_point(&self) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
let operations = std::mem::take(&mut inner.operations);
inner.savepoints.push(Savepoint {
operations,
});
Ok(())
})
}
fn release_last_save_point(&self) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
inner.savepoints.pop();
Ok(())
})
}
fn rollback_to_save_point(&self) -> BoxFut<'_, Result<()>> {
Box::pin(async move {
if self.closed() {
return Err(Error::TransactionFinished);
}
if !self.writeable() {
return Err(Error::TransactionReadonly);
}
let mut inner = self.inner.write().await;
if inner.savepoints.is_empty() {
return Err(Error::Transaction("No savepoint to rollback to".to_string()));
}
let savepoint = inner.savepoints.pop().expect("No savepoint to rollback to");
let operations = std::mem::take(&mut inner.operations);
for op in operations.iter().rev() {
match op {
Operation::DeleteKey(key) => {
inner.tx.delete(key.clone()).await?;
}
Operation::RestoreValue(key, val) => {
inner.tx.put(key.clone(), val.clone()).await?;
}
Operation::RestoreDeleted(key, val) => {
inner.tx.put(key.clone(), val.clone()).await?;
}
}
}
inner.operations = savepoint.operations;
Ok(())
})
}
fn timestamp(&self) -> BoxFut<'_, Result<BoxTimeStamp>> {
Box::pin(async move {
let ts = self.inner.write().await.tx.current_timestamp().await?;
Ok(BoxTimeStamp::new(TiKVStamp(ts)))
})
}
fn timestamp_impl(&self) -> BoxTimeStampImpl {
Box::new(TiKVStampImpl)
}
}
pub struct TiKVStampImpl;
impl TimeStampImpl for TiKVStampImpl {
fn earliest(&self) -> BoxTimeStamp {
BoxTimeStamp::new(TiKVStamp(tikv::Timestamp {
physical: 0,
logical: 0,
suffix_bits: 0,
}))
}
fn create_from_versionstamp(&self, version: u128) -> Option<BoxTimeStamp> {
Some(BoxTimeStamp::new(TiKVStamp(tikv::Timestamp::from_version(version as u64))))
}
fn create_from_datetime(&self, dt: DateTime<Utc>) -> Option<BoxTimeStamp> {
let physical = dt.timestamp_millis();
Some(BoxTimeStamp::new(TiKVStamp(tikv::Timestamp {
physical,
logical: 0,
suffix_bits: 0,
})))
}
fn decode(&self, bytes: &[u8]) -> Result<BoxTimeStamp> {
if bytes.len() == 8 {
let Ok(b) = <[u8; 8]>::try_from(&bytes[0..8]) else {
unreachable!()
};
let ts = u64::from_be_bytes(b);
return Ok(BoxTimeStamp::new(TiKVStamp(tikv::Timestamp::from_version(ts))));
}
if bytes.len() != 20 {
return Err(Error::TimestampInvalid(
"Encoded timestamp is not the right length".to_string(),
));
}
let Ok(b) = <[u8; 8]>::try_from(&bytes[0..8]) else {
unreachable!()
};
let physical = i64::from_be_bytes(b) ^ i64::MIN;
let Ok(b) = <[u8; 8]>::try_from(&bytes[8..16]) else {
unreachable!()
};
let logical = i64::from_be_bytes(b) ^ i64::MIN;
let Ok(b) = <[u8; 4]>::try_from(&bytes[16..20]) else {
unreachable!()
};
let suffix_bits = u32::from_be_bytes(b);
Ok(BoxTimeStamp::new(TiKVStamp(tikv::Timestamp {
physical,
logical,
suffix_bits,
})))
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TiKVStamp(tikv::Timestamp);
impl TimeStamp for TiKVStamp {
fn as_versionstamp(&self) -> u128 {
self.0.version() as u128
}
fn as_datetime(&self) -> Option<DateTime<Utc>> {
DateTime::from_timestamp_millis(self.0.physical)
}
fn sub_checked(&self, duration: Duration) -> Option<BoxTimeStamp> {
let millis = duration.as_millis().try_into().ok()?;
let physical = self.0.physical.checked_sub(millis)?;
Some(BoxTimeStamp::new(TiKVStamp(tikv::Timestamp {
physical,
logical: self.0.logical,
suffix_bits: self.0.suffix_bits,
})))
}
fn encode<'a>(&self, bytes: &'a mut [u8; MAX_TIMESTAMP_BYTES]) -> &'a [u8] {
let b = (self.0.physical ^ i64::MIN).to_be_bytes();
bytes[0..8].copy_from_slice(&b);
let b = (self.0.logical ^ i64::MIN).to_be_bytes();
bytes[8..16].copy_from_slice(&b);
let b = self.0.suffix_bits.to_be_bytes();
bytes[16..20].copy_from_slice(&b);
&bytes[..20]
}
}
fn consume_keys<I: Iterator<Item = tikv::Key>>(iter: &mut I, limit: u32, skip: u32) -> KeysResult {
for _ in 0..skip {
if iter.next().is_none() {
return KeysResult::default();
}
}
let mut key_bytes = 0u64;
let mut keys = Vec::with_capacity(limit.min(4096) as usize);
while keys.len() < limit as usize {
if let Some(k) = iter.next() {
key_bytes += k.len() as u64;
keys.push(Key::from(k));
} else {
break;
}
}
KeysResult {
keys,
key_bytes,
}
}
fn consume_vals<I: Iterator<Item = tikv::KvPair>>(iter: &mut I, limit: u32) -> ScanResult {
let mut key_bytes = 0u64;
let mut value_bytes = 0u64;
let mut values = Vec::with_capacity(limit.min(4096) as usize);
while values.len() < limit as usize {
if let Some(kv) = iter.next() {
let key_len = kv.0.len() as u64;
let value_len = kv.1.len() as u64;
key_bytes += key_len;
value_bytes += value_len;
values.push((Key::from(kv.0), kv.1));
} else {
break;
}
}
ScanResult {
values,
key_bytes,
value_bytes,
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use super::{TikvConfig, safepoint_from, validate_tls_paths};
#[test]
fn safepoint_from_subtracts_lifetime() {
let now = tikv::Timestamp {
physical: 1_000_000,
logical: 42,
suffix_bits: 7,
};
let safepoint = safepoint_from(&now, Duration::from_millis(250)).unwrap();
assert_eq!(safepoint.physical, 999_750);
assert_eq!(safepoint.logical, 42);
assert_eq!(safepoint.suffix_bits, 7);
}
#[test]
fn safepoint_from_returns_none_when_lifetime_too_large() {
let now = tikv::Timestamp {
physical: 100,
logical: 0,
suffix_bits: 0,
};
assert!(safepoint_from(&now, Duration::from_millis(101)).is_none());
}
#[test]
fn safepoint_from_returns_none_when_lifetime_overflows_i64() {
let now = tikv::Timestamp {
physical: i64::MAX,
logical: 0,
suffix_bits: 0,
};
assert!(safepoint_from(&now, Duration::MAX).is_none());
}
#[test]
fn validate_tls_paths_all_unset() {
let cfg = TikvConfig::default();
assert!(validate_tls_paths(&cfg).unwrap().is_none());
}
#[test]
fn validate_tls_paths_all_set() {
let cfg = TikvConfig {
tls_ca_path: Some("/etc/tikv/ca.pem".into()),
tls_cert_path: Some("/etc/tikv/cert.pem".into()),
tls_key_path: Some("/etc/tikv/key.pem".into()),
..TikvConfig::default()
};
let paths = validate_tls_paths(&cfg).unwrap().unwrap();
assert_eq!(paths.ca, "/etc/tikv/ca.pem");
assert_eq!(paths.cert, "/etc/tikv/cert.pem");
assert_eq!(paths.key, "/etc/tikv/key.pem");
}
#[test]
fn validate_tls_paths_partial_set_rejects() {
let cases: &[(Option<&str>, Option<&str>, Option<&str>)] = &[
(Some("ca"), None, None),
(None, Some("cert"), None),
(None, None, Some("key")),
(Some("ca"), Some("cert"), None),
(Some("ca"), None, Some("key")),
(None, Some("cert"), Some("key")),
];
for (ca, cert, key) in cases {
let cfg = TikvConfig {
tls_ca_path: ca.map(String::from),
tls_cert_path: cert.map(String::from),
tls_key_path: key.map(String::from),
..TikvConfig::default()
};
let err = validate_tls_paths(&cfg).expect_err("partial TLS config must reject");
let msg = err.to_string();
assert!(
msg.contains("tikv_tls_ca_path"),
"error should name the env var triple: {msg}"
);
assert!(msg.contains("set"), "error should report which fields are set: {msg}");
assert!(msg.contains("unset"), "error should report which fields are unset: {msg}");
}
}
#[test]
fn release_in_flight_counter_decrements() {
let counter = Arc::new(AtomicUsize::new(2));
let release = |c: &Arc<AtomicUsize>| {
let _ = c.fetch_update(Ordering::AcqRel, Ordering::Acquire, |n| n.checked_sub(1));
};
release(&counter);
assert_eq!(counter.load(Ordering::Acquire), 1);
release(&counter);
assert_eq!(counter.load(Ordering::Acquire), 0);
release(&counter);
assert_eq!(counter.load(Ordering::Acquire), 0);
}
}