use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use crate::core::{StoreError, TenantId, Timestamp};
use crate::journal::payload;
use crate::memory::{Cascade, MemoryItem, MemoryStore, Recall, Selected};
use super::{Erasure, KeyError, KeyRing};
#[derive(Debug, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct PlainMemory {
content: serde_json::Value,
derived_from: Vec<Selected>,
}
pub struct EncryptedMemoryStore {
inner: Arc<dyn MemoryStore>,
keys: Arc<dyn KeyRing>,
tenant: TenantId,
lifecycle: Arc<dyn super::ErasureCoordinator>,
}
impl std::fmt::Debug for EncryptedMemoryStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EncryptedMemoryStore")
.field("tenant", &self.tenant)
.finish_non_exhaustive()
}
}
impl EncryptedMemoryStore {
#[must_use]
pub fn new(inner: Arc<dyn MemoryStore>, keys: Arc<dyn KeyRing>, tenant: TenantId) -> Self {
super::assert_serves(inner.tenant(), &tenant, "memory");
Self {
inner,
keys,
tenant,
lifecycle: Arc::new(super::LocalCoordinator::new()),
}
}
#[must_use]
pub fn coordinated_by(mut self, coordinator: Arc<dyn super::ErasureCoordinator>) -> Self {
self.lifecycle = coordinator;
self
}
#[must_use]
pub fn is_distributed(&self) -> bool {
self.lifecycle.is_distributed()
}
fn lifecycle_scope(&self) -> String {
super::scope(&self.tenant, "memory-lifecycle")
}
fn scope(&self, id: &str, version: u64) -> String {
super::scope(&self.tenant, &format!("memory-item/{id}@{version}"))
}
fn aad(&self, item: &MemoryItem, version: u64) -> Result<Vec<u8>, StoreError> {
crate::core::canon::to_bytes(&(
"memory",
self.tenant.as_str(),
item.id.as_str(),
version,
item.subject.as_str(),
item.purpose.as_str(),
))
.map_err(|error| StoreError::Backend(error.to_string()))
}
async fn seal(&self, item: &MemoryItem, version: u64) -> Result<serde_json::Value, StoreError> {
let plain = crate::core::canon::to_bytes(&PlainMemory {
content: item.content.clone(),
derived_from: item.derived_from.clone(),
})
.map_err(|error| StoreError::Backend(error.to_string()))?;
let envelope = super::envelope::seal(
self.keys.as_ref(),
&self.scope(&item.id, version),
&self.aad(item, version)?,
&plain,
)
.await
.map_err(|error| match error {
KeyError::Destroyed { .. } => StoreError::Backend(format!(
"memory id '{}' was erased and cannot be reused",
item.id
)),
other => key_error(other),
})?;
Ok(payload::wrap(&envelope))
}
async fn open_item(&self, mut item: MemoryItem) -> Result<Option<MemoryItem>, StoreError> {
let envelope = payload::unwrap(&item.content).ok_or_else(|| {
StoreError::Backend(
"encrypted memory row does not contain a sealed envelope".to_owned(),
)
})?;
let aad = self.aad(&item, item.version)?;
let Some(plain) = super::envelope::open_or_erased(self.keys.as_ref(), &aad, &envelope)
.await
.map_err(key_error)?
else {
return Ok(None);
};
let plain: PlainMemory = serde_json::from_slice(&plain)
.map_err(|error| StoreError::Backend(format!("encrypted memory: {error}")))?;
item.content = plain.content;
item.derived_from = plain.derived_from;
Ok(Some(item))
}
async fn backing_selection(&self, source: &Selected) -> Result<Selected, StoreError> {
let stored = self
.inner
.version(&source.id, source.version)
.await?
.ok_or_else(|| {
StoreError::Backend(format!(
"derived memory source '{}' version {} is absent",
source.id, source.version
))
})?;
let opened = self.open_item(stored.clone()).await?.ok_or_else(|| {
StoreError::Backend(format!(
"derived memory source '{}' version {} was erased",
source.id, source.version
))
})?;
if opened.selection_digest() != source.digest {
return Err(StoreError::Backend(format!(
"derived memory source '{}' version {} changed",
source.id, source.version
)));
}
Ok(Selected {
id: source.id.clone(),
version: source.version,
digest: stored.selection_digest(),
})
}
async fn destroy_erased(
&self,
erased: &[(String, Vec<u64>)],
at: Timestamp,
reason: &str,
) -> Result<(), StoreError> {
for (id, versions) in erased {
for version in versions {
let scope = self.scope(id, *version);
self.keys
.destroy(&scope, at, reason)
.await
.map_err(|error| {
StoreError::Backend(format!(
"memory '{id}' version {version} was erased from the store, and \
destroying its key failed ({error}) — its backups still open until \
scope '{scope}' is destroyed"
))
})?;
}
}
Ok(())
}
fn every_version(ids: impl IntoIterator<Item = (String, u64)>) -> Vec<(String, Vec<u64>)> {
ids.into_iter()
.map(|(id, highest)| (id, (1..=highest).collect()))
.collect()
}
async fn highest_versions(&self, ids: &[String]) -> Result<Vec<(String, u64)>, StoreError> {
let mut out = Vec::with_capacity(ids.len());
for id in ids {
if let Some(current) = self.inner.current(id, None).await? {
out.push((id.clone(), current.version));
}
}
Ok(out)
}
async fn refuse_held(&self, ids: &[String]) -> Result<(), StoreError> {
for id in ids {
if self.inner.legal_hold(id).await? {
return Err(StoreError::UnderLegalHold { id: id.clone() });
}
}
Ok(())
}
pub async fn erase_subject(
&self,
subject: &str,
at: Timestamp,
reason: &str,
) -> Result<Erasure, StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
let ids = self.inner.subject_ids(subject).await?;
self.refuse_held(&ids).await?;
for (id, versions) in Self::every_version(self.highest_versions(&ids).await?) {
for version in versions {
self.keys
.destroy(&self.scope(&id, version), at, reason)
.await
.map_err(key_error)?;
}
}
Ok(match self.inner.forget_subject(subject).await {
Ok(count) => Erasure {
reached: count,
cleanup_failed: None,
},
Err(error) => {
tracing::warn!(%subject, %error, "memory keys were destroyed but ciphertext cleanup failed");
Erasure {
reached: ids.len(),
cleanup_failed: Some(error.to_string()),
}
}
})
})
.await
}
}
#[allow(clippy::needless_pass_by_value)]
fn key_error(error: KeyError) -> StoreError {
StoreError::Backend(error.to_string())
}
#[allow(clippy::disallowed_methods)]
fn verb_erasure(verb: &str) -> (Timestamp, String) {
(Timestamp::now_utc(), format!("memory {verb}"))
}
#[async_trait]
impl MemoryStore for EncryptedMemoryStore {
fn tenant(&self) -> &str {
self.tenant.as_str()
}
fn erasure_is_distributed(&self) -> Option<bool> {
Some(self.lifecycle.is_distributed())
}
async fn remember(&self, item: &MemoryItem) -> Result<u64, StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
let version = self
.inner
.current(&item.id, None)
.await?
.map_or(1, |current| current.version + 1);
let mut sealed = item.clone();
sealed.content = self.seal(item, version).await?;
sealed.derived_from.clear();
for source in &item.derived_from {
sealed
.derived_from
.push(self.backing_selection(source).await?);
}
let written = self.inner.remember(&sealed).await?;
if written != version {
return Err(StoreError::Backend(format!(
"memory '{}' was sealed as version {version} and stored as {written}; \
the row will not open, so the write is refused",
item.id
)));
}
Ok(written)
})
.await
}
async fn recall(&self, query: &Recall) -> Result<Vec<MemoryItem>, StoreError> {
let items = self.inner.recall(query).await?;
let mut opened = Vec::with_capacity(items.len());
for item in items {
if let Some(item) = self.open_item(item).await? {
opened.push(item);
}
}
Ok(opened)
}
async fn subject_ids(&self, subject: &str) -> Result<Vec<String>, StoreError> {
self.inner.subject_ids(subject).await
}
async fn version(&self, id: &str, version: u64) -> Result<Option<MemoryItem>, StoreError> {
match self.inner.version(id, version).await? {
Some(item) => self.open_item(item).await,
None => Ok(None),
}
}
async fn current(
&self,
id: &str,
as_of: Option<Timestamp>,
) -> Result<Option<MemoryItem>, StoreError> {
match self.inner.current(id, as_of).await? {
Some(item) => self.open_item(item).await,
None => Ok(None),
}
}
async fn forget(&self, id: &str) -> Result<(), StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
self.refuse_held(&[id.to_owned()]).await?;
let (at, reason) = verb_erasure("forget");
for (id, versions) in
Self::every_version(self.highest_versions(&[id.to_owned()]).await?)
{
for version in versions {
self.keys
.destroy(&self.scope(&id, version), at, &reason)
.await
.map_err(key_error)?;
}
}
self.inner.forget(id).await.map_err(|error| {
StoreError::Backend(format!(
"memory '{id}': its key is destroyed, so no copy opens, and removing its \
rows failed ({error}) — retry forget to remove them"
))
})
})
.await
}
async fn forget_subject(&self, subject: &str) -> Result<usize, StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
let ids = self.inner.subject_ids(subject).await?;
let highest = self.highest_versions(&ids).await?;
let count = self.inner.forget_subject(subject).await?;
let (at, reason) = verb_erasure("forget_subject");
self.destroy_erased(&Self::every_version(highest), at, &reason)
.await?;
Ok(count)
})
.await
}
async fn derivatives(&self, id: &str) -> Result<Vec<MemoryItem>, StoreError> {
let items = self.inner.derivatives(id).await?;
let mut opened = Vec::with_capacity(items.len());
for item in items {
if let Some(item) = self.open_item(item).await? {
opened.push(item);
}
}
Ok(opened)
}
async fn forget_cascading(&self, id: &str) -> Result<Cascade, StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
let cascade = self.inner.forget_cascading(id).await?;
let (at, reason) = verb_erasure("forget_cascading");
self.destroy_erased(&Self::every_version(cascade.erased.clone()), at, &reason)
.await?;
self.destroy_erased(&cascade.trimmed, at, &reason).await?;
Ok(cascade)
})
.await
}
async fn set_legal_hold(&self, id: &str, held: bool) -> Result<(), StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
self.inner.set_legal_hold(id, held).await
})
.await
}
async fn legal_holds(
&self,
after: Option<&str>,
limit: usize,
) -> Result<Vec<String>, StoreError> {
self.inner.legal_holds(after, limit).await
}
async fn legal_hold(&self, id: &str) -> Result<bool, StoreError> {
self.inner.legal_hold(id).await
}
async fn sweep_expired(&self, at: Timestamp) -> Result<Vec<(String, u64)>, StoreError> {
super::under_lock(self.lifecycle.as_ref(), &self.lifecycle_scope(), || async {
let swept = self.inner.sweep_expired(at).await?;
self.destroy_erased(
&Self::every_version(swept.clone()),
at,
"memory retention expired",
)
.await?;
Ok(swept)
})
.await
}
async fn touch(&self, ids: &[String], at: Timestamp) -> Result<(), StoreError> {
self.inner.touch(ids, at).await
}
}