use std::collections::VecDeque;
use std::sync::{Arc, Mutex, PoisonError};
use sha2::{Digest, Sha256};
use crate::types::Message;
const BOUNDARY_RING_CAPACITY: usize = 8;
#[derive(Debug, Clone)]
struct Midstate {
hasher: Sha256,
covered: usize,
}
impl Midstate {
fn finalize(&self) -> String {
let mut hasher = self.hasher.clone();
hasher.update(b"]");
let digest = hasher.finalize();
let mut out = String::with_capacity(digest.len() * 2 + 7);
out.push_str("sha256:");
const HEX: &[u8; 16] = b"0123456789abcdef";
for byte in digest {
out.push(HEX[(byte >> 4) as usize] as char);
out.push(HEX[(byte & 0x0f) as usize] as char);
}
out
}
fn absorb(&mut self, message: &Message) -> Result<(), serde_json::Error> {
if self.covered > 0 {
self.hasher.update(b",");
}
let canonical = super::canonicalize_message_for_digest(message);
let bytes = serde_json::to_vec(&canonical)?;
crate::checkpoint::record_content_digest_bytes(bytes.len() as u64);
self.hasher.update(bytes);
self.covered += 1;
Ok(())
}
}
#[derive(Debug, Clone)]
struct SortedStream {
bytes: Vec<u8>,
covered: usize,
}
#[derive(Debug, Clone, Default)]
struct AccumulatorState {
stream_a: Option<Midstate>,
boundaries: VecDeque<Midstate>,
sorted_stream: Option<SortedStream>,
epoch: u64,
parked: Option<Box<AccumulatorState>>,
}
#[derive(Debug, Default)]
pub(crate) struct TranscriptDigestAccumulator {
state: Mutex<Box<AccumulatorState>>,
}
impl Clone for TranscriptDigestAccumulator {
fn clone(&self) -> Self {
Self {
state: Mutex::new(self.locked().clone()),
}
}
}
thread_local! {
static CROSS_CHECK_ENABLED: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
}
static RELEASE_VERIFICATION_BUDGET: std::sync::atomic::AtomicU32 =
std::sync::atomic::AtomicU32::new(32);
pub(crate) fn take_verification_sample() -> bool {
if cfg!(any(test, debug_assertions)) {
CROSS_CHECK_ENABLED.with(std::cell::Cell::get)
} else {
RELEASE_VERIFICATION_BUDGET
.fetch_update(
std::sync::atomic::Ordering::Relaxed,
std::sync::atomic::Ordering::Relaxed,
|budget| budget.checked_sub(1),
)
.is_ok()
}
}
impl TranscriptDigestAccumulator {
fn locked(&self) -> std::sync::MutexGuard<'_, Box<AccumulatorState>> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn epoch(&self) -> u64 {
self.locked().epoch
}
fn invalidate(&mut self) {
let state = self.state.get_mut().unwrap_or_else(PoisonError::into_inner);
state.stream_a = None;
state.boundaries.clear();
state.sorted_stream = None;
state.parked = None;
state.epoch = state.epoch.saturating_add(1);
}
fn extend(&mut self, appended: &[Message]) {
let state = self.state.get_mut().unwrap_or_else(PoisonError::into_inner);
if let Some(stream) = state.stream_a.as_mut() {
for message in appended {
if stream.absorb(message).is_err() {
state.stream_a = None;
state.boundaries.clear();
state.sorted_stream = None;
state.epoch = state.epoch.saturating_add(1);
return;
}
}
}
if let Some(sorted) = state.sorted_stream.as_mut() {
for message in appended {
match sorted_canonical_message_bytes(message) {
Ok(bytes) => {
if sorted.covered > 0 {
sorted.bytes.push(b',');
}
sorted.bytes.extend_from_slice(&bytes);
sorted.covered += 1;
}
Err(_) => {
state.sorted_stream = None;
break;
}
}
}
}
}
fn begin_in_place_scan(&mut self) {
let state = self.state.get_mut().unwrap_or_else(PoisonError::into_inner);
if state.stream_a.is_none() && state.boundaries.is_empty() && state.sorted_stream.is_none()
{
return;
}
let parked = AccumulatorState {
stream_a: state.stream_a.take(),
boundaries: std::mem::take(&mut state.boundaries),
sorted_stream: state.sorted_stream.take(),
epoch: state.epoch,
parked: None,
};
state.parked = Some(Box::new(parked));
}
fn finish_in_place_scan(&mut self, lowest_mutated_index: Option<usize>) {
let state = self.state.get_mut().unwrap_or_else(PoisonError::into_inner);
let Some(parked) = state.parked.take() else {
if lowest_mutated_index.is_some() {
state.stream_a = None;
state.boundaries.clear();
state.epoch = state.epoch.saturating_add(1);
}
return;
};
if lowest_mutated_index.is_some() {
state.epoch = state.epoch.saturating_add(1);
return;
}
state.stream_a = parked.stream_a;
state.boundaries = parked.boundaries;
state.sorted_stream = parked.sorted_stream;
}
fn digest(&self, messages: &[Message]) -> Result<String, serde_json::Error> {
if let Some(witness) = self.witness(messages) {
let mut state = self.locked();
if let Some(stream) = state.stream_a.clone() {
record_boundary(&mut state, &stream);
}
return Ok(witness);
}
let mut state = self.locked();
let mut stream = Midstate {
hasher: Sha256::new(),
covered: 0,
};
stream.hasher.update(b"[");
crate::checkpoint::record_content_digest_computation();
for message in messages {
stream.absorb(message)?;
}
let digest = stream.finalize();
record_boundary(&mut state, &stream);
state.stream_a = Some(stream);
Ok(digest)
}
fn witness(&self, messages: &[Message]) -> Option<String> {
let state = self.locked();
let stream = state.stream_a.as_ref()?;
if stream.covered != messages.len() {
return None;
}
let digest = stream.finalize();
drop(state);
if take_verification_sample()
&& let Ok(recomputed) = super::transcript_messages_digest_uncounted(messages)
{
assert_eq!(
digest, recomputed,
"transcript digest accumulator served a stale witness: a message-mutation \
seam extended or replaced the transcript without invalidating the midstate"
);
}
Some(digest)
}
fn prefix_witness(&self, messages: &[Message], count: usize) -> Option<String> {
if count > messages.len() {
return None;
}
let state = self.locked();
let boundary = state
.boundaries
.iter()
.find(|midstate| midstate.covered == count)
.or_else(|| {
state
.stream_a
.as_ref()
.filter(|stream| stream.covered == count)
})?;
let digest = boundary.finalize();
drop(state);
if take_verification_sample()
&& let Ok(recomputed) = super::transcript_messages_digest_uncounted(&messages[..count])
{
assert_eq!(
digest, recomputed,
"transcript digest accumulator served a stale prefix witness: a \
message-mutation seam rewrote a retained prefix without invalidating the ring"
);
}
Some(digest)
}
fn hash_sorted_canonical_into(&self, messages: &[Message], hasher: &mut Sha256) -> bool {
let mut state = self.locked();
if state.parked.is_some() {
return false;
}
let serve = match state.sorted_stream.as_ref() {
Some(sorted) if sorted.covered == messages.len() => true,
_ => {
let mut sorted = SortedStream {
bytes: Vec::new(),
covered: 0,
};
for message in messages {
let Ok(bytes) = sorted_canonical_message_bytes(message) else {
return false;
};
if sorted.covered > 0 {
sorted.bytes.push(b',');
}
sorted.bytes.extend_from_slice(&bytes);
sorted.covered += 1;
}
state.sorted_stream = Some(sorted);
true
}
};
if !serve {
return false;
}
let Some(sorted) = state.sorted_stream.as_ref() else {
return false;
};
hasher.update(b"[");
hasher.update(&sorted.bytes);
hasher.update(b"]");
crate::checkpoint::record_content_digest_bytes(sorted.bytes.len() as u64 + 2);
true
}
}
fn sorted_canonical_message_bytes(message: &Message) -> Result<Vec<u8>, serde_json::Error> {
let canonical = serde_json::to_value(super::canonicalize_message_for_digest(message))?;
let mut bytes = Vec::new();
crate::checkpoint::write_canonical_json(&canonical, &mut bytes)?;
Ok(bytes)
}
fn record_boundary(state: &mut AccumulatorState, stream: &Midstate) {
if state
.boundaries
.iter()
.any(|midstate| midstate.covered == stream.covered)
{
return;
}
if state.boundaries.len() >= BOUNDARY_RING_CAPACITY {
state.boundaries.pop_front();
}
state.boundaries.push_back(stream.clone());
}
#[derive(Debug, Default)]
pub(crate) struct TranscriptMessages {
messages: Arc<Vec<Message>>,
accumulator: Box<TranscriptDigestAccumulator>,
}
impl Clone for TranscriptMessages {
fn clone(&self) -> Self {
Self {
messages: Arc::clone(&self.messages),
accumulator: Box::new((*self.accumulator).clone()),
}
}
}
impl std::ops::Deref for TranscriptMessages {
type Target = Vec<Message>;
fn deref(&self) -> &Self::Target {
&self.messages
}
}
impl TranscriptMessages {
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn arc(&self) -> &Arc<Vec<Message>> {
&self.messages
}
pub(crate) fn from_vec(messages: Vec<Message>) -> Self {
Self {
messages: Arc::new(messages),
accumulator: Box::default(),
}
}
pub(crate) fn push(&mut self, message: Message) {
let Self {
messages,
accumulator,
} = self;
let inner = Arc::make_mut(messages);
inner.push(message);
let appended = &inner[inner.len() - 1..];
accumulator.extend(appended);
}
pub(crate) fn extend_batch(&mut self, appended: Vec<Message>) {
if appended.is_empty() {
return;
}
let Self {
messages,
accumulator,
} = self;
let inner = Arc::make_mut(messages);
let start = inner.len();
inner.extend(appended);
accumulator.extend(&inner[start..]);
}
pub(crate) fn replace(&mut self, messages: Vec<Message>) {
self.messages = Arc::new(messages);
self.accumulator.invalidate();
}
pub(crate) fn mutate_in_place(&mut self) -> &mut Vec<Message> {
self.accumulator.invalidate();
Arc::make_mut(&mut self.messages)
}
pub(crate) fn begin_in_place_scan(&mut self) -> &mut Vec<Message> {
self.accumulator.begin_in_place_scan();
Arc::make_mut(&mut self.messages)
}
pub(crate) fn finish_in_place_scan(&mut self, lowest_mutated_index: Option<usize>) {
self.accumulator.finish_in_place_scan(lowest_mutated_index);
}
pub(crate) fn digest(&self) -> Result<String, serde_json::Error> {
self.accumulator.digest(&self.messages)
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn digest_witness(&self) -> Option<String> {
self.accumulator.witness(&self.messages)
}
pub(crate) fn prefix_digest_witness(&self, count: usize) -> Option<String> {
self.accumulator.prefix_witness(&self.messages, count)
}
pub(crate) fn mutation_epoch(&self) -> u64 {
self.accumulator.epoch()
}
pub(crate) fn hash_sorted_canonical_into(&self, hasher: &mut Sha256) -> bool {
self.accumulator
.hash_sorted_canonical_into(&self.messages, hasher)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use crate::types::{Message, UserMessage};
fn user(text: &str) -> Message {
Message::User(UserMessage::text(text))
}
fn transcript(count: usize) -> Vec<Message> {
(0..count).map(|i| user(&format!("m{i}"))).collect()
}
#[test]
fn canonicalize_messages_for_digest_is_element_wise() {
let messages = transcript(6);
let whole = super::super::canonicalize_messages_for_digest(&messages);
let per_message = messages
.iter()
.map(super::super::canonicalize_message_for_digest)
.collect::<Vec<_>>();
assert_eq!(whole, per_message);
let array_bytes = serde_json::to_vec(&whole).unwrap();
let mut streamed = Vec::from(b"[".as_slice());
for (index, message) in per_message.iter().enumerate() {
if index > 0 {
streamed.extend_from_slice(b",");
}
streamed.extend_from_slice(&serde_json::to_vec(message).unwrap());
}
streamed.extend_from_slice(b"]");
assert_eq!(array_bytes, streamed);
}
#[test]
fn seeded_digest_matches_full_recompute() {
for count in [0usize, 1, 2, 7, 40] {
let messages = TranscriptMessages::from_vec(transcript(count));
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap(),
"count {count}"
);
}
}
#[test]
fn appended_digest_matches_full_recompute() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let _ = messages.digest().unwrap();
messages.push(user("appended"));
assert!(messages.digest_witness().is_some());
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
messages.extend_batch(vec![user("a"), user("b")]);
assert!(messages.digest_witness().is_some());
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
}
#[test]
fn unseeded_accumulator_has_no_witness() {
let messages = TranscriptMessages::from_vec(transcript(3));
assert!(messages.digest_witness().is_none());
assert!(messages.prefix_digest_witness(2).is_none());
}
#[test]
fn replacement_invalidates_the_witness() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let _ = messages.digest().unwrap();
let epoch = messages.mutation_epoch();
messages.replace(transcript(2));
assert!(messages.digest_witness().is_none());
assert!(messages.mutation_epoch() > epoch);
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
}
#[test]
fn in_place_mutation_invalidates_the_witness() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let _ = messages.digest().unwrap();
messages.mutate_in_place()[0] = user("rewritten");
assert!(messages.digest_witness().is_none());
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
}
#[test]
fn unmutated_in_place_scan_keeps_the_witness() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let seeded = messages.digest().unwrap();
let _buffer = messages.begin_in_place_scan();
assert!(
messages.digest_witness().is_none(),
"a parked accumulator must not serve a witness"
);
messages.finish_in_place_scan(None);
assert_eq!(messages.digest_witness(), Some(seeded));
}
#[test]
fn mutated_in_place_scan_drops_the_witness() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let _ = messages.digest().unwrap();
{
let buffer = messages.begin_in_place_scan();
buffer[1] = user("changed");
}
messages.finish_in_place_scan(Some(1));
assert!(messages.digest_witness().is_none());
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
}
#[test]
fn abandoned_in_place_scan_fails_safe() {
let mut messages = TranscriptMessages::from_vec(transcript(3));
let _ = messages.digest().unwrap();
let buffer = messages.begin_in_place_scan();
buffer[2] = user("changed");
assert!(messages.digest_witness().is_none());
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap()
);
}
#[test]
fn boundary_ring_answers_prefix_queries_after_appends() {
let mut messages = TranscriptMessages::from_vec(transcript(5));
let boundary = messages.digest().unwrap();
messages.extend_batch(vec![user("x"), user("y")]);
assert_eq!(messages.prefix_digest_witness(5), Some(boundary));
assert_eq!(
messages.prefix_digest_witness(5).unwrap(),
super::super::transcript_messages_digest(&messages[..5]).unwrap()
);
assert!(messages.prefix_digest_witness(4).is_none());
}
#[test]
fn boundary_ring_is_bounded_and_dropped_on_invalidation() {
let mut messages = TranscriptMessages::from_vec(Vec::new());
for _ in 0..(BOUNDARY_RING_CAPACITY + 4) {
messages.push(user("m"));
let _ = messages.digest().unwrap();
}
let retained = (0..=(BOUNDARY_RING_CAPACITY + 4))
.filter(|count| messages.prefix_digest_witness(*count).is_some())
.count();
assert!(
retained <= BOUNDARY_RING_CAPACITY,
"boundary ring grew past its bound: {retained}"
);
messages.mutate_in_place();
assert_eq!(
(0..=(BOUNDARY_RING_CAPACITY + 4))
.filter(|count| messages.prefix_digest_witness(*count).is_some())
.count(),
0
);
}
#[test]
fn randomized_mutation_sequences_match_full_recompute() {
let mut seed = 0x5eed_1234_u64;
let mut next = move || {
seed = seed
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(seed >> 33) as usize
};
let mut messages = TranscriptMessages::from_vec(transcript(4));
for step in 0..200 {
match next() % 6 {
0 => messages.push(user(&format!("p{step}"))),
1 => messages.extend_batch(vec![user(&format!("b{step}")), user("b2")]),
2 => messages.replace(transcript(next() % 9)),
3 => {
let buffer = messages.mutate_in_place();
if !buffer.is_empty() {
let index = next() % buffer.len();
buffer[index] = user(&format!("r{step}"));
}
}
4 => {
messages.begin_in_place_scan();
messages.finish_in_place_scan(None);
}
_ => {
let buffer = messages.begin_in_place_scan();
let mutated = if buffer.is_empty() {
None
} else {
let index = next() % buffer.len();
buffer[index] = user(&format!("s{step}"));
Some(index)
};
messages.finish_in_place_scan(mutated);
}
}
assert_eq!(
messages.digest().unwrap(),
super::super::transcript_messages_digest(&messages).unwrap(),
"step {step}"
);
let count = messages.len();
if count > 1 {
messages.push(user("tail"));
assert_eq!(
messages.prefix_digest_witness(count),
Some(super::super::transcript_messages_digest(&messages[..count]).unwrap()),
"prefix witness at step {step}"
);
}
}
}
}