use std::collections::HashMap;
use std::fmt;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::{Mutex, PoisonError};
use super::{Capacity, Row, SqlConn, SqlError, SqlValue, TxFn, blob, count, schema};
use crate::store::{
Batch, BatchOutcome, Cursor, Key, NamespaceStore, Partition, PartitionStats, Precondition,
ScanPage, StoreCapabilities, StoreError, StoreMaintenance, Value, Write, codec, keys,
};
pub const GET_MANY_CHUNK: usize = 64;
const STATS_TTL_MS: u64 = 60_000;
const STATS_CACHE_MAX: usize = 4096;
pub(super) const GET: &str = "SELECT value FROM kv WHERE part = ?1 AND key = ?2";
pub(super) const PUT: &str = "INSERT INTO kv (part, key, value) VALUES (?1, ?2, ?3) \
ON CONFLICT (part, key) DO UPDATE SET value = excluded.value";
pub(super) const DELETE: &str = "DELETE FROM kv WHERE part = ?1 AND key = ?2";
pub(super) const SCAN_FROM: &str = "SELECT key, value FROM kv \
WHERE part = ?1 AND key >= ?2 AND key < ?3 ORDER BY key LIMIT ?4";
pub(super) const SCAN_AFTER: &str = "SELECT key, value FROM kv \
WHERE part = ?1 AND key > ?2 AND key < ?3 ORDER BY key LIMIT ?4";
pub(super) const STATS: &str =
"SELECT COUNT(*), SUM(length(key) + length(value)) FROM kv WHERE part = ?1";
pub const TIMER_WINDOW_START: &str = "SELECT part, key FROM kv INDEXED BY kv_timers \
WHERE key >= x'7700' AND key < x'7701' ORDER BY key, part LIMIT ?1";
pub const TIMER_WINDOW_AFTER: &str = "SELECT part, key FROM kv INDEXED BY kv_timers \
WHERE (key, part) > (?1, ?2) AND key >= x'7700' AND key < x'7701' \
ORDER BY key, part LIMIT ?3";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TimerCursor {
pub key: Key,
pub partition: Partition,
}
pub(super) const PROBE: &str = "SELECT 1";
pub(super) fn get_many_sql(n: usize) -> String {
let marks: Vec<String> = (2..n + 2).map(|i| format!("?{i}")).collect();
format!(
"SELECT key, value FROM kv WHERE part = ?1 AND key IN ({})",
marks.join(", ")
)
}
pub struct SqlKvStore<C> {
conn: C,
schema_version: AtomicU32,
stats: Mutex<HashMap<Vec<u8>, (u64, PartitionStats)>>,
capacity: Option<Capacity>,
}
impl<C> fmt::Debug for SqlKvStore<C> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SqlKvStore")
.field(
"schema_version",
&self.schema_version.load(Ordering::Relaxed),
)
.field("capacity", &self.capacity)
.finish_non_exhaustive()
}
}
impl<C: SqlConn> SqlKvStore<C> {
pub fn open(conn: C) -> Result<Self, StoreError> {
let version = schema::migrate(&conn)?;
Ok(Self::from_checked_conn(conn, version))
}
pub fn open_existing(conn: C) -> Result<Self, StoreError> {
let version = schema::require_current(&conn)?;
Ok(Self::from_checked_conn(conn, version))
}
fn from_checked_conn(conn: C, version: u32) -> Self {
Self {
conn,
schema_version: AtomicU32::new(version),
stats: Mutex::default(),
capacity: None,
}
}
pub fn open_with_capacity(conn: C, capacity: Capacity) -> Result<Self, StoreError> {
let mut store = Self::open(conn)?;
store.conn.set_size_limit(capacity.cap_bytes())?;
store.capacity = Some(capacity);
Ok(store)
}
pub fn apply_extended<F>(
&self,
p: &Partition,
batch: Batch,
extend: F,
) -> Result<BatchOutcome, StoreError>
where
F: FnOnce(
&dyn Fn(&Key) -> Result<Option<Value>, SqlError>,
&mut Batch,
u64,
) -> Result<(), SqlError>
+ 'static,
{
batch.validate(&self.capabilities())?;
let part = part(p)?;
let adds = batch.has_put();
let soft_limit = self.capacity.map(|c| c.soft_limit());
match self.conn.transaction(Box::new(move |conn| {
check_and_write_extended(&conn, &part, batch, soft_limit, extend)
})) {
Ok(outcome) => Ok(outcome),
Err(SqlError::Full) if !adds => Err(StoreError::unavailable(
"database full during a delete-only batch",
)),
Err(error) => Err(error.into()),
}
}
pub fn timer_window(
&self,
after: Option<&TimerCursor>,
limit: u32,
) -> Result<Vec<TimerCursor>, StoreError> {
if limit == 0 {
return Err(StoreError::Invalid(
"timer window limit must be positive".into(),
));
}
let (sql, params) = match after {
None => (
TIMER_WINDOW_START,
vec![SqlValue::Integer(i64::from(limit))],
),
Some(cursor) => (
TIMER_WINDOW_AFTER,
vec![
SqlValue::Blob(cursor.key.as_bytes().to_vec()),
SqlValue::Blob(cursor.partition.encode()?.to_vec()),
SqlValue::Integer(i64::from(limit)),
],
),
};
self.conn
.query(sql, ¶ms)?
.into_iter()
.map(|mut row| {
Ok(TimerCursor {
partition: Partition::decode(&blob(&mut row, 0)?)?,
key: Key::new(blob(&mut row, 1)?),
})
})
.collect()
}
#[must_use]
pub fn capacity(&self) -> Option<Capacity> {
self.capacity
}
#[must_use]
pub fn conn(&self) -> &C {
&self.conn
}
pub fn clear_stats_cache(&self) {
self.stats
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clear();
}
}
fn part(p: &Partition) -> Result<SqlValue, StoreError> {
Ok(SqlValue::Blob(p.encode()?.to_vec()))
}
fn key_param(key: &Key) -> SqlValue {
SqlValue::Blob(key.as_bytes().to_vec())
}
fn read<C: SqlConn>(conn: &C, part: &SqlValue, key: &Key) -> Result<Option<Value>, SqlError> {
let mut rows = conn.query(GET, &[part.clone(), key_param(key)])?;
rows.first_mut()
.map(|row| blob(row, 0).map(Value::new))
.transpose()
}
fn check_and_write<C: SqlConn>(
conn: &C,
part: &SqlValue,
batch: Batch,
soft_limit: Option<u64>,
) -> Result<BatchOutcome, SqlError> {
check_and_write_extended(conn, part, batch, soft_limit, |_, _, _| Ok(()))
}
fn check_and_write_extended<C: SqlConn, F>(
conn: &C,
part: &SqlValue,
mut batch: Batch,
soft_limit: Option<u64>,
extend: F,
) -> Result<BatchOutcome, SqlError>
where
F: FnOnce(
&dyn Fn(&Key) -> Result<Option<Value>, SqlError>,
&mut Batch,
u64,
) -> Result<(), SqlError>,
{
let now = conn.now_ms();
for (index, pre) in batch.preconditions.iter().enumerate() {
let (holds, observed) = match pre {
Precondition::NotAfter(deadline) if now > *deadline => {
return Ok(BatchOutcome::DeadlinePassed { backend_now: now });
}
Precondition::NotAfter(_) => continue,
Precondition::Absent(key) => {
let current = read(conn, part, key)?;
(current.is_none(), current)
}
Precondition::Present(key) => (read(conn, part, key)?.is_some(), None),
Precondition::Equals(key, value) => {
let current = read(conn, part, key)?;
(current.as_ref() == Some(value), current)
}
};
if !holds {
return Ok(BatchOutcome::PreconditionFailed { index, observed });
}
}
extend(&|key| read(conn, part, key), &mut batch, now)?;
batch
.validate(&StoreCapabilities::full())
.map_err(|_| SqlError::Corrupt("extended batch exceeds limits"))?;
if let Some(limit) = soft_limit
&& batch.has_put()
&& !is_relay_scan_checkpoint(&batch)
&& !is_relay_timer_reschedule(&batch)
&& !is_timer_retry_move(&batch)
&& conn.size_bytes()? >= limit
{
return Err(SqlError::Full);
}
for write in batch.writes {
match write {
Write::Put(key, value) => {
let value = SqlValue::Blob(value.into_bytes().into());
conn.exec(PUT, &[part.clone(), key_param(&key), value])?
}
Write::Delete(key) => conn.exec(DELETE, &[part.clone(), key_param(&key)])?,
};
}
Ok(BatchOutcome::Committed)
}
fn is_relay_scan_checkpoint(batch: &Batch) -> bool {
let scan_key = keys::relay_scan();
let guarded = batch.preconditions.iter().any(|pre| {
matches!(pre, Precondition::Absent(key) | Precondition::Equals(key, _) if key == &scan_key)
});
let mut scan_puts = 0;
let only_cleanup = batch.writes.iter().all(|write| match write {
Write::Put(key, value) if key == &scan_key => {
scan_puts += 1;
codec::decode_relay_scan(value).is_ok()
}
Write::Delete(key) => matches!(keys::parse(key), Some(keys::ParsedKey::Relay(_))),
Write::Put(..) => false,
});
guarded && only_cleanup && scan_puts == 1
}
pub(super) fn is_relay_timer_reschedule(batch: &Batch) -> bool {
let (
[
Precondition::Equals(old_key, old_value),
Precondition::Absent(absent),
],
[Write::Delete(deleted), Write::Put(new_key, new_value)],
) = (batch.preconditions.as_slice(), batch.writes.as_slice())
else {
return false;
};
if old_key != deleted
|| absent != new_key
|| !old_value.as_bytes().is_empty()
|| !new_value.as_bytes().is_empty()
{
return false;
}
let (
Some(keys::ParsedKey::Timer {
due_at_ms: old_due,
kind: old_kind,
reference: old_ref,
}),
Some(keys::ParsedKey::Timer {
due_at_ms: new_due,
kind: new_kind,
reference: new_ref,
}),
) = (keys::parse(old_key), keys::parse(new_key))
else {
return false;
};
old_kind == crate::timers::registry::kinds::RELAY.get()
&& new_kind == old_kind
&& old_ref == new_ref
&& new_due > old_due
}
pub(super) fn is_timer_retry_move(batch: &Batch) -> bool {
let (
[
Precondition::Equals(old_key, old_value),
Precondition::Absent(absent),
],
[Write::Delete(deleted), Write::Put(new_key, new_value)],
) = (batch.preconditions.as_slice(), batch.writes.as_slice())
else {
return false;
};
if old_key != deleted || absent != new_key || old_value != new_value {
return false;
}
let (
Some(keys::ParsedKey::Timer {
due_at_ms: old_due,
kind: old_kind,
reference: old_ref,
}),
Some(keys::ParsedKey::Timer {
due_at_ms: new_due,
kind: new_kind,
reference: new_ref,
}),
Some((old_original, old_attempt)),
Some((new_original, new_attempt)),
) = (
keys::parse(old_key),
keys::parse(new_key),
keys::timer_retry_state(old_key),
keys::timer_retry_state(new_key),
)
else {
return false;
};
new_due > old_due
&& new_kind == old_kind
&& new_ref == old_ref
&& new_original == old_original
&& new_attempt
== old_attempt
.saturating_add(1)
.min(keys::MAX_TIMER_RETRY_ATTEMPT)
}
fn entry(mut row: Row) -> Result<(Key, Value), SqlError> {
Ok((Key::new(blob(&mut row, 0)?), Value::new(blob(&mut row, 1)?)))
}
impl<C: SqlConn> NamespaceStore for SqlKvStore<C> {
fn capabilities(&self) -> StoreCapabilities {
StoreCapabilities::full()
}
async fn get(&self, p: &Partition, key: &Key) -> Result<Option<Value>, StoreError> {
Ok(read(&self.conn, &part(p)?, key)?)
}
async fn get_many(
&self,
p: &Partition,
keys: &[Key],
) -> Result<Vec<Option<Value>>, StoreError> {
let part = part(p)?;
let mut found: HashMap<Vec<u8>, Value> = HashMap::new();
for chunk in keys.chunks(GET_MANY_CHUNK) {
let params: Vec<SqlValue> = core::iter::once(part.clone())
.chain(chunk.iter().map(key_param))
.collect();
for row in self.conn.query(&get_many_sql(chunk.len()), ¶ms)? {
let (key, value) = entry(row)?;
found.insert(key.into_bytes().into(), value);
}
}
Ok(keys
.iter()
.map(|key| found.get(key.as_bytes()).cloned())
.collect())
}
async fn scan(
&self,
p: &Partition,
start: &Key,
end: &Key,
after: Option<&Cursor>,
limit: u32,
) -> Result<ScanPage, StoreError> {
if limit == 0 {
return Err(StoreError::Invalid("scan limit must be at least 1".into()));
}
let (sql, lower) = match after {
None => (SCAN_FROM, start.as_bytes()),
Some(c) if start.as_bytes() <= c.as_bytes() && c.as_bytes() < end.as_bytes() => {
(SCAN_AFTER, c.as_bytes())
}
Some(_) => {
return Err(StoreError::Invalid(
"scan cursor outside the scanned range".into(),
));
}
};
if lower >= end.as_bytes() {
return Ok(ScanPage::default());
}
let params = [
part(p)?,
SqlValue::Blob(lower.to_vec()),
key_param(end),
SqlValue::Integer(i64::from(limit) + 1),
];
let mut entries = self
.conn
.query(sql, ¶ms)?
.into_iter()
.map(entry)
.collect::<Result<Vec<_>, _>>()?;
let want = usize::try_from(limit).unwrap_or(usize::MAX);
let next = (entries.len() > want).then(|| {
entries.truncate(want);
Cursor::new(entries[want - 1].0.clone().into_bytes())
});
Ok(ScanPage { entries, next })
}
async fn apply(&self, p: &Partition, batch: Batch) -> Result<BatchOutcome, StoreError> {
batch.validate(&self.capabilities())?;
let part = part(p)?;
let adds = batch.has_put();
let soft_limit = self.capacity.map(|c| c.soft_limit());
let body: TxFn<C, BatchOutcome> =
Box::new(move |conn: C| check_and_write(&conn, &part, batch, soft_limit));
match self.conn.transaction(body) {
Ok(outcome) => Ok(outcome),
Err(SqlError::Full) if !adds => Err(StoreError::unavailable(
"database full during a delete-only batch",
)),
Err(e) => Err(e.into()),
}
}
async fn stats(&self, p: &Partition) -> Result<PartitionStats, StoreError> {
let name = p.encode()?.to_vec();
let now = self.conn.now_ms();
let fresh = |at: u64| at <= now && now - at < STATS_TTL_MS;
{
let cache = self.stats.lock().unwrap_or_else(PoisonError::into_inner);
if let Some((at, stats)) = cache.get(&name)
&& fresh(*at)
{
return Ok(*stats);
}
}
let rows = self.conn.query(STATS, &[SqlValue::Blob(name.clone())])?;
let row = rows
.first()
.ok_or(SqlError::Corrupt("stats returned no row"))?;
let stats = PartitionStats {
bytes: count(row, 1)?,
keys: Some(count(row, 0)?),
};
let mut cache = self.stats.lock().unwrap_or_else(PoisonError::into_inner);
if cache.len() >= STATS_CACHE_MAX {
cache.retain(|_, (at, _)| fresh(*at));
if cache.len() >= STATS_CACHE_MAX {
cache.clear();
}
}
cache.insert(name, (now, stats));
Ok(stats)
}
async fn probe(&self) -> Result<(), StoreError> {
self.conn.query(PROBE, &[])?;
Ok(())
}
}
impl<C: SqlConn> StoreMaintenance for SqlKvStore<C> {
fn layout_version(&self) -> u32 {
self.schema_version.load(Ordering::Relaxed)
}
async fn migrate(&self) -> Result<u32, StoreError> {
let version = schema::migrate(&self.conn)?;
self.schema_version.store(version, Ordering::Relaxed);
Ok(version)
}
async fn backup_to(&self, dest: &str) -> Result<(), StoreError> {
self.conn.backup_to(dest)
}
}