use std::marker::PhantomData;
use std::sync::Arc;
use serde::{de::DeserializeOwned, Serialize};
use zeroize::Zeroizing;
use crate::codec::{decode_value, encode_value, KeyDecode, KeyEncode};
use crate::compress::{pack_value, unpack_value, Compression};
use crate::crypto::{open_row, seal_row, unwrap_dek, wrap_fresh_dek, KeyProvider, WrappedDek, KEY_LEN};
use crate::engine::{Durability, KvEngine, Readable, WriteTx};
use crate::error::{StoreError, StoreResult};
use crate::table::Table;
pub const DEK_WRAPPINGS_TABLE: &str = "_dek_wrappings";
pub const COLLECTION_FORMATS_TABLE: &str = "_collection_formats";
const VALUE_FORMAT_PREFIXED: u8 = 2;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ValueFormat {
Legacy,
Prefixed,
}
fn wrappings_table() -> Table<String, WrappedDek> {
Table::new(DEK_WRAPPINGS_TABLE)
}
fn formats_table() -> Table<String, u8> {
Table::new(COLLECTION_FORMATS_TABLE)
}
fn read_format(tx: &impl Readable, name: &str) -> StoreResult<ValueFormat> {
Ok(match formats_table().get(tx, &name.to_string())? {
Some(VALUE_FORMAT_PREFIXED) => ValueFormat::Prefixed,
_ => ValueFormat::Legacy,
})
}
pub struct Collection<K, V> {
name: String,
schema_version: u32,
wrapped_dek: WrappedDek,
key_provider: Arc<dyn KeyProvider>,
format: ValueFormat,
compression: Compression,
_types: PhantomData<fn() -> (K, V)>,
}
impl<K, V> std::fmt::Debug for Collection<K, V> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Collection")
.field("name", &self.name)
.field("schema_version", &self.schema_version)
.finish_non_exhaustive()
}
}
impl<K, V> Collection<K, V>
where
K: KeyEncode + KeyDecode,
V: Serialize + DeserializeOwned,
{
pub fn open<E: KvEngine>(
engine: &E,
key_provider: Arc<dyn KeyProvider>,
name: impl Into<String>,
schema_version: u32,
) -> StoreResult<Self> {
Self::open_with(engine, key_provider, name, schema_version, Compression::default())
}
pub fn open_with<E: KvEngine>(
engine: &E,
key_provider: Arc<dyn KeyProvider>,
name: impl Into<String>,
schema_version: u32,
compression: Compression,
) -> StoreResult<Self> {
let name = Self::checked_name(name)?;
let r = engine.begin_read()?;
let wrapped = wrappings_table()
.get(&r, &name)?
.ok_or_else(|| StoreError::CollectionNotFound(name.clone()))?;
let format = read_format(&r, &name)?;
Ok(Self::assemble(name, schema_version, wrapped, key_provider, format, compression))
}
pub fn open_or_create<E: KvEngine>(
engine: &E,
key_provider: Arc<dyn KeyProvider>,
name: impl Into<String>,
schema_version: u32,
) -> StoreResult<Self> {
Self::open_or_create_with(engine, key_provider, name, schema_version, Compression::default())
}
pub fn open_or_create_with<E: KvEngine>(
engine: &E,
key_provider: Arc<dyn KeyProvider>,
name: impl Into<String>,
schema_version: u32,
compression: Compression,
) -> StoreResult<Self> {
let name = Self::checked_name(name)?;
let table = wrappings_table();
let mut w = engine.begin_write(Durability::Immediate)?;
if let Some(existing) = table.get(&w, &name)? {
let format = read_format(&w, &name)?;
drop(w);
return Ok(Self::assemble(name, schema_version, existing, key_provider, format, compression));
}
let vault_key = key_provider.vault_key()?;
let (wrapped, _dek) = wrap_fresh_dek(&vault_key, &name)?;
table.put(&mut w, &name, &wrapped)?;
formats_table().put(&mut w, &name, &VALUE_FORMAT_PREFIXED)?;
w.commit()?;
Ok(Self::assemble(
name,
schema_version,
wrapped,
key_provider,
ValueFormat::Prefixed,
compression,
))
}
fn assemble(
name: String,
schema_version: u32,
wrapped_dek: WrappedDek,
key_provider: Arc<dyn KeyProvider>,
format: ValueFormat,
compression: Compression,
) -> Self {
Self {
name,
schema_version,
wrapped_dek,
key_provider,
format,
compression,
_types: PhantomData,
}
}
fn checked_name(name: impl Into<String>) -> StoreResult<String> {
let name = name.into();
if name.starts_with('_') {
return Err(StoreError::ReservedName(name));
}
Ok(name)
}
pub fn name(&self) -> &str {
&self.name
}
pub fn schema_version(&self) -> u32 {
self.schema_version
}
fn dek(&self) -> StoreResult<Zeroizing<[u8; KEY_LEN]>> {
let vault_key = self.key_provider.vault_key()?;
Ok(unwrap_dek(&vault_key, &self.name, &self.wrapped_dek)?)
}
fn plaintext_for_store(&self, value: &V) -> StoreResult<Vec<u8>> {
let plain = encode_value(value)?;
Ok(match self.format {
ValueFormat::Legacy => plain,
ValueFormat::Prefixed => pack_value(&plain, self.compression),
})
}
fn value_from_plaintext(&self, plain: &[u8]) -> StoreResult<V> {
match self.format {
ValueFormat::Legacy => decode_value(plain),
ValueFormat::Prefixed => decode_value(&unpack_value(plain)?),
}
}
pub fn get(&self, tx: &impl Readable, key: &K) -> StoreResult<Option<V>> {
let key_bytes = key.encode();
let sealed = match tx.get_raw(&self.name, &key_bytes)? {
Some(bytes) => bytes,
None => return Ok(None),
};
let dek = self.dek()?;
let plain = open_row(&dek, &self.name, &key_bytes, self.schema_version, &sealed)?;
Ok(Some(self.value_from_plaintext(&plain)?))
}
pub fn put(&self, tx: &mut impl WriteTx, key: &K, value: &V) -> StoreResult<()> {
let key_bytes = key.encode();
let dek = self.dek()?;
let sealed = seal_row(
&dek,
&self.name,
&key_bytes,
self.schema_version,
&self.plaintext_for_store(value)?,
)?;
tx.put_raw(&self.name, &key_bytes, &sealed)
}
pub fn delete(&self, tx: &mut impl WriteTx, key: &K) -> StoreResult<bool> {
tx.delete_raw(&self.name, &key.encode())
}
pub fn range(&self, tx: &impl Readable, lo: &K, hi: &K) -> StoreResult<Vec<(K, V)>> {
let raw = tx.range_raw(&self.name, &lo.encode(), &hi.encode())?;
if raw.is_empty() {
return Ok(Vec::new());
}
let dek = self.dek()?;
raw.into_iter()
.map(|(key_bytes, sealed)| {
let plain = open_row(&dek, &self.name, &key_bytes, self.schema_version, &sealed)?;
Ok((K::decode(&key_bytes)?, self.value_from_plaintext(&plain)?))
})
.collect()
}
}