use rabs_protocol::result_identity::{DigestAlgorithm, TypedDigest};
use sha2::{Digest, Sha256};
pub const ATP_OBJECT_CONTENT_DOMAIN: &str = "rabs.object.sha256.v1";
pub const BLAKE3_FINGERPRINT_DOMAIN: &str = "rabs.fingerprint.blake3.v1";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Blake3Fingerprint {
pub domain: &'static str,
pub bytes: [u8; 32],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct RawSha256(pub [u8; 32]);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DigestSet {
pub atp_content_id: TypedDigest,
pub blake3: Option<Blake3Fingerprint>,
pub raw_sha256: Option<RawSha256>,
pub logical_size: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct DigestRequest {
pub blake3: bool,
pub raw_sha256: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DigestError {
LogicalSizeMismatch {
expected: u64,
actual: u64,
},
SizeOverflow,
}
fn domain_prefix(domain: &str) -> ([u8; 8], &[u8]) {
((domain.len() as u64).to_be_bytes(), domain.as_bytes())
}
pub struct StreamingObjectWriter {
content: Sha256,
blake3: Option<blake3::Hasher>,
raw: Option<Sha256>,
expected_size: Option<u64>,
written: u64,
}
impl StreamingObjectWriter {
#[must_use]
pub fn new(request: DigestRequest, expected_size: Option<u64>) -> Self {
let mut content = Sha256::new();
let (len, domain) = domain_prefix(ATP_OBJECT_CONTENT_DOMAIN);
content.update(len);
content.update(domain);
let blake3 = request.blake3.then(|| {
let mut hasher = blake3::Hasher::new();
let (len, domain) = domain_prefix(BLAKE3_FINGERPRINT_DOMAIN);
hasher.update(&len);
hasher.update(domain);
hasher
});
let raw = request.raw_sha256.then(Sha256::new);
Self {
content,
blake3,
raw,
expected_size,
written: 0,
}
}
pub fn write(&mut self, chunk: &[u8]) -> Result<(), DigestError> {
self.written = self
.written
.checked_add(chunk.len() as u64)
.ok_or(DigestError::SizeOverflow)?;
self.content.update(chunk);
if let Some(hasher) = &mut self.blake3 {
hasher.update(chunk);
}
if let Some(hasher) = &mut self.raw {
hasher.update(chunk);
}
Ok(())
}
pub fn finish(self) -> Result<DigestSet, DigestError> {
if let Some(expected) = self.expected_size
&& expected != self.written
{
return Err(DigestError::LogicalSizeMismatch {
expected,
actual: self.written,
});
}
Ok(DigestSet {
atp_content_id: TypedDigest {
algorithm: DigestAlgorithm::Sha256V1,
domain: ATP_OBJECT_CONTENT_DOMAIN,
bytes: self.content.finalize().into(),
},
blake3: self.blake3.map(|hasher| Blake3Fingerprint {
domain: BLAKE3_FINGERPRINT_DOMAIN,
bytes: *hasher.finalize().as_bytes(),
}),
raw_sha256: self.raw.map(|hasher| RawSha256(hasher.finalize().into())),
logical_size: self.written,
})
}
}
pub fn digest_set(
bytes: &[u8],
request: DigestRequest,
expected_size: Option<u64>,
) -> Result<DigestSet, DigestError> {
let mut writer = StreamingObjectWriter::new(request, expected_size);
writer.write(bytes)?;
writer.finish()
}
#[cfg(test)]
mod tests {
use super::*;
const ALL: DigestRequest = DigestRequest {
blake3: true,
raw_sha256: true,
};
#[test]
fn h002_streaming_equals_oneshot_under_arbitrary_chunking() {
let bytes: Vec<u8> = (0..u8::MAX).cycle().take(70_001).collect();
let oneshot = digest_set(&bytes, ALL, Some(bytes.len() as u64)).unwrap();
for splits in [
vec![0usize],
vec![1, 2, 3],
vec![70_000],
vec![16 * 1024, 40_000, 69_999, 70_000],
] {
let mut writer = StreamingObjectWriter::new(ALL, Some(bytes.len() as u64));
let mut last = 0;
for split in splits {
writer.write(&bytes[last..split]).unwrap();
last = split;
}
writer.write(&bytes[last..]).unwrap();
assert_eq!(writer.finish().unwrap(), oneshot);
}
assert_eq!(oneshot.logical_size, 70_001);
}
#[test]
fn h002_tags_bind_algorithm_and_domain_per_role() {
let set = digest_set(b"tag enforcement", ALL, None).unwrap();
assert_eq!(set.atp_content_id.algorithm, DigestAlgorithm::Sha256V1);
assert_eq!(set.atp_content_id.domain, ATP_OBJECT_CONTENT_DOMAIN);
let blake = set.blake3.unwrap();
assert_eq!(blake.domain, BLAKE3_FINGERPRINT_DOMAIN);
let raw = set.raw_sha256.unwrap();
assert_ne!(set.atp_content_id.bytes, blake.bytes);
assert_ne!(set.atp_content_id.bytes, raw.0);
assert_ne!(blake.bytes, raw.0);
let mut reference = Sha256::new();
reference.update((ATP_OBJECT_CONTENT_DOMAIN.len() as u64).to_be_bytes());
reference.update(ATP_OBJECT_CONTENT_DOMAIN.as_bytes());
reference.update(b"tag enforcement");
assert_eq!(
set.atp_content_id.bytes,
<[u8; 32]>::from(reference.finalize())
);
let mut reference = blake3::Hasher::new();
reference.update(&(BLAKE3_FINGERPRINT_DOMAIN.len() as u64).to_be_bytes());
reference.update(BLAKE3_FINGERPRINT_DOMAIN.as_bytes());
reference.update(b"tag enforcement");
assert_eq!(blake.bytes, *reference.finalize().as_bytes());
let mut reference = Sha256::new();
reference.update(b"tag enforcement");
assert_eq!(raw.0, <[u8; 32]>::from(reference.finalize()));
}
#[test]
fn h002_raw_sha256_matches_published_sha2_test_vector() {
let set = digest_set(
b"abc",
DigestRequest {
blake3: false,
raw_sha256: true,
},
None,
)
.unwrap();
let expected: [u8; 32] = [
0xba, 0x78, 0x16, 0xbf, 0x8f, 0x01, 0xcf, 0xea, 0x41, 0x41, 0x40, 0xde, 0x5d, 0xae,
0x22, 0x23, 0xb0, 0x03, 0x61, 0xa3, 0x96, 0x17, 0x7a, 0x9c, 0xb4, 0x10, 0xff, 0x61,
0xf2, 0x00, 0x15, 0xad,
];
assert_eq!(set.raw_sha256.unwrap().0, expected);
assert!(set.blake3.is_none(), "unrequested digests are not computed");
}
#[test]
fn h002_logical_size_verification_refuses_short_and_long_writes() {
let mut writer = StreamingObjectWriter::new(DigestRequest::default(), Some(4));
writer.write(b"abc").unwrap();
assert_eq!(
writer.finish(),
Err(DigestError::LogicalSizeMismatch {
expected: 4,
actual: 3
})
);
let mut writer = StreamingObjectWriter::new(DigestRequest::default(), Some(2));
writer.write(b"abc").unwrap();
assert_eq!(
writer.finish(),
Err(DigestError::LogicalSizeMismatch {
expected: 2,
actual: 3
})
);
let mut writer = StreamingObjectWriter::new(DigestRequest::default(), Some(3));
writer.write(b"abc").unwrap();
assert_eq!(writer.finish().unwrap().logical_size, 3);
}
#[test]
fn h002_encoding_is_separate_from_identity() {
let logical = b"logical object bytes, eminently compressible aaaaaaaaaaaaaaaa";
let pretend_zstd: Vec<u8> = logical.iter().rev().copied().collect();
let logical_set = digest_set(logical, DigestRequest::default(), None).unwrap();
let encoded_set = digest_set(&pretend_zstd, DigestRequest::default(), None).unwrap();
assert_ne!(logical_set.atp_content_id, encoded_set.atp_content_id);
assert_eq!(
logical_set.atp_content_id.domain,
encoded_set.atp_content_id.domain
);
}
}