use std::sync::Arc;
use async_trait::async_trait;
use serde_json::Value;
use crate::case::{CaseStore, Correlation};
use crate::core::{
Case, CaseId, CaseStatus, CaseVersion, CorrelationKey, Deadline, DeadlineState, Digest,
LegalHold, RunId, StoreError, TenantId, Timestamp,
};
use crate::journal::payload;
use super::KeyRing;
#[derive(Debug)]
pub struct SealedCases {
inner: Arc<dyn CaseStore>,
keys: Arc<dyn KeyRing>,
tenant: TenantId,
}
impl SealedCases {
#[must_use]
pub fn wrap(inner: Arc<dyn CaseStore>, keys: Arc<dyn KeyRing>, tenant: TenantId) -> Arc<Self> {
super::assert_serves(inner.tenant(), &tenant, "case");
Arc::new(Self {
inner,
keys,
tenant,
})
}
fn scope_for(&self, case: CaseId) -> String {
super::scope(&self.tenant, &case.to_string())
}
fn aad(&self, case: CaseId) -> String {
format!("case-state:{}:{case}", self.tenant)
}
async fn open_state(&self, case: CaseId, state: Value) -> Result<Value, StoreError> {
let Some(envelope) = payload::unwrap(&state) else {
return Ok(state);
};
let aad = self.aad(case);
let opened = super::envelope::open_or_erased(self.keys.as_ref(), aad.as_bytes(), &envelope)
.await
.map_err(|e| StoreError::Backend(e.to_string()))?;
Ok(match opened {
Some(plain) => serde_json::from_slice(&plain).unwrap_or(state),
None => state,
})
}
async fn open_case(&self, mut case: Case) -> Result<Case, StoreError> {
case.state = self
.open_state(case.id, std::mem::take(&mut case.state))
.await?;
Ok(case)
}
async fn opened(&self, case: Option<Case>) -> Result<Option<Case>, StoreError> {
Ok(match case {
Some(case) => Some(self.open_case(case).await?),
None => None,
})
}
}
pub async fn probe_sealed_case_state(
keys: &dyn KeyRing,
case: CaseId,
state: &Value,
) -> Option<Result<(), super::KeyError>> {
if !payload::is_sealed(state) {
return None;
}
let Some(envelope) = payload::unwrap(state) else {
return Some(Err(super::KeyError::Refused(
"the state is marked sealed and its envelope is not valid base64".to_owned(),
)));
};
let scope = match super::envelope::wrapped_scope(&envelope) {
Ok(scope) => scope,
Err(e) => return Some(Err(e)),
};
let Some(tenant) = scope.strip_suffix(&format!("/{case}")) else {
return Some(Err(super::KeyError::Refused(format!(
"this case's sealed state names erasure scope '{scope}', which is not this \
case — the envelope was written for a different matter, so erasing this \
case would leave it readable"
))));
};
let aad = format!("case-state:{tenant}:{case}");
Some(
super::envelope::open(keys, aad.as_bytes(), &envelope)
.await
.map(drop),
)
}
#[async_trait]
impl CaseStore for SealedCases {
fn tenant(&self) -> &str {
self.tenant.as_str()
}
async fn put_state(
&self,
case: CaseId,
expected: CaseVersion,
state: Value,
) -> Result<CaseVersion, StoreError> {
let plain = crate::core::canon::to_bytes(&state)
.map_err(|e| StoreError::Backend(format!("case state would not serialise: {e}")))?;
let envelope = super::envelope::seal(
self.keys.as_ref(),
&self.scope_for(case),
self.aad(case).as_bytes(),
&plain,
)
.await
.map_err(|e| StoreError::Backend(format!("sealing case state failed: {e}")))?;
self.inner
.put_state(case, expected, payload::wrap(&envelope))
.await
}
async fn case(&self, id: CaseId) -> Result<Option<Case>, StoreError> {
let found = self.inner.case(id).await?;
self.opened(found).await
}
async fn by_status(&self, status: CaseStatus, limit: usize) -> Result<Vec<Case>, StoreError> {
let cases = self.inner.by_status(status, limit).await?;
let mut out = Vec::with_capacity(cases.len());
for case in cases {
out.push(self.open_case(case).await?);
}
Ok(out)
}
async fn correlate(&self, keys: &[CorrelationKey]) -> Result<Option<CaseId>, StoreError> {
self.inner.correlate(keys).await
}
async fn correlate_or_open(
&self,
kind: &str,
keys: &[CorrelationKey],
at: Timestamp,
) -> Result<Correlation, StoreError> {
self.inner.correlate_or_open(kind, keys, at).await
}
async fn cases(&self, after: Option<CaseId>, limit: usize) -> Result<Vec<Case>, StoreError> {
self.inner.cases(after, limit).await
}
async fn import_case(
&self,
case: &Case,
deadlines: &[crate::core::Deadline],
blobs: &[crate::core::Digest],
) -> Result<(), StoreError> {
self.inner.import_case(case, deadlines, blobs).await
}
async fn attach_run(&self, case: CaseId, run: RunId) -> Result<(), StoreError> {
self.inner.attach_run(case, run).await
}
async fn detach_run(&self, case: CaseId, run: RunId) -> Result<bool, StoreError> {
self.inner.detach_run(case, run).await
}
async fn link_blob(
&self,
case: CaseId,
digest: Digest,
at: Timestamp,
) -> Result<(), StoreError> {
self.inner.link_blob(case, digest, at).await
}
async fn blobs_of(&self, case: CaseId) -> Result<Vec<Digest>, StoreError> {
self.inner.blobs_of(case).await
}
async fn set_status(&self, case: CaseId, status: CaseStatus) -> Result<(), StoreError> {
self.inner.set_status(case, status).await
}
async fn close(&self, case: CaseId) -> Result<(), StoreError> {
self.inner.close(case).await
}
async fn place_hold(&self, case: CaseId, hold: &LegalHold) -> Result<bool, StoreError> {
self.inner.place_hold(case, hold).await
}
async fn release_hold(&self, case: CaseId) -> Result<bool, StoreError> {
self.inner.release_hold(case).await
}
async fn hold(&self, case: CaseId) -> Result<Option<LegalHold>, StoreError> {
self.inner.hold(case).await
}
async fn holds(
&self,
after: Option<CaseId>,
limit: usize,
) -> Result<Vec<(CaseId, LegalHold)>, StoreError> {
self.inner.holds(after, limit).await
}
async fn register_deadline(&self, deadline: &Deadline) -> Result<(), StoreError> {
self.inner.register_deadline(deadline).await
}
async fn deadlines(&self, case: CaseId) -> Result<Vec<Deadline>, StoreError> {
self.inner.deadlines(case).await
}
async fn set_deadline_state(
&self,
case: CaseId,
name: &str,
state: DeadlineState,
) -> Result<(), StoreError> {
self.inner.set_deadline_state(case, name, state).await
}
async fn breach_deadline(
&self,
case: CaseId,
name: &str,
now: Timestamp,
) -> Result<bool, StoreError> {
self.inner.breach_deadline(case, name, now).await
}
async fn due(&self, now: Timestamp, limit: usize) -> Result<Vec<Deadline>, StoreError> {
self.inner.due(now, limit).await
}
async fn breaches_to_note(&self, limit: usize) -> Result<Vec<Deadline>, StoreError> {
self.inner.breaches_to_note(limit).await
}
async fn mark_breach_noted(&self, case: CaseId, name: &str) -> Result<(), StoreError> {
self.inner.mark_breach_noted(case, name).await
}
async fn breached(&self, limit: usize) -> Result<Vec<Deadline>, StoreError> {
self.inner.breached(limit).await
}
async fn acknowledge_breach(
&self,
case: CaseId,
name: &str,
note: &crate::core::BreachNote,
) -> Result<bool, StoreError> {
self.inner.acknowledge_breach(case, name, note).await
}
async fn census(&self, now: Timestamp) -> Result<crate::case::CaseCensus, StoreError> {
self.inner.census(now).await
}
async fn record_drill(&self, record: &crate::case::DrillRecord) -> Result<(), StoreError> {
self.inner.record_drill(record).await
}
async fn last_drill(&self) -> Result<Option<crate::case::DrillRecord>, StoreError> {
self.inner.last_drill().await
}
}
#[cfg(all(test, feature = "testkit"))]
mod probe_tests {
use super::*;
use crate::testkit::MemoryKeyRing;
use serde_json::json;
fn matter() -> CaseId {
CaseId::generate()
}
async fn sealed_state(ring: &MemoryKeyRing, tenant: &str, case: CaseId) -> Value {
let aad = format!("case-state:{tenant}:{case}");
let plain = crate::core::canon::to_bytes(&json!({ "about": "the matter" })).expect("canon");
let envelope =
super::super::envelope::seal(ring, &format!("{tenant}/{case}"), aad.as_bytes(), &plain)
.await
.expect("seal");
payload::wrap(&envelope)
}
#[tokio::test]
async fn state_that_was_never_sealed_is_an_ordinary_absence() {
let ring = MemoryKeyRing::new();
assert!(
probe_sealed_case_state(&ring, matter(), &json!({ "about": "plain" }))
.await
.is_none(),
"unsealed state is not something a sealing probe has an answer about"
);
}
#[tokio::test]
async fn sealed_state_that_opens_reports_that_it_opened() {
let ring = MemoryKeyRing::new();
let case = matter();
let state = sealed_state(&ring, "acme", case).await;
assert_eq!(
probe_sealed_case_state(&ring, case, &state).await,
Some(Ok(())),
);
}
#[tokio::test]
async fn sealed_state_this_build_cannot_read_is_answered_not_skipped() {
let ring = MemoryKeyRing::new();
let case = matter();
let state = sealed_state(&ring, "acme", case).await;
let envelope = payload::unwrap(&state).expect("sealed");
let mut bumped = envelope.clone();
bumped[0] = bumped[0].wrapping_add(1);
let rows: [(&str, CaseId, Value); 3] = [
(
"a format version this build does not read",
case,
payload::wrap(&bumped),
),
(
"an envelope sealed for another matter",
matter(),
state.clone(),
),
(
"a marker whose payload is not base64",
case,
json!({ payload::SEALED: "not base64 !!" }),
),
];
for (label, case, value) in rows {
assert!(
payload::is_sealed(&value),
"{label}: the row stopped claiming to be sealed, so it proves nothing"
);
let answer = probe_sealed_case_state(&ring, case, &value).await;
assert!(
matches!(answer, Some(Err(_))),
"{label} was reported as nothing to check: {answer:?}"
);
}
assert!(
matches!(
probe_sealed_case_state(&ring, case, &payload::wrap(&bumped)).await,
Some(Err(super::super::KeyError::UnknownFormat { .. }))
),
"a version skew must not reach the drill as a suspected loss"
);
}
}