use super::SceneDescription;
use crate::Scene;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::io::{self, Read, Write};
use thiserror::Error;
const ADDRESS_BYTES: usize = 32;
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug, Serialize, Deserialize)]
pub struct ContentAddress([u8; ADDRESS_BYTES]);
impl ContentAddress {
#[must_use]
pub fn digest(bytes: &[u8]) -> Self {
let mut hasher = Sha256::new();
hasher.update(bytes);
Self(hasher.finalize().into())
}
#[must_use]
pub const fn bytes(self) -> [u8; ADDRESS_BYTES] {
self.0
}
}
pub(crate) struct ContentHasher(Sha256);
impl ContentHasher {
pub(crate) fn new() -> Self {
Self(Sha256::new())
}
pub(crate) fn update(&mut self, bytes: &[u8]) {
self.0.update(bytes);
}
pub(crate) fn finish(self) -> ContentAddress {
ContentAddress(self.0.finalize().into())
}
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub enum ReferencedPayloadKind {
Structure,
Property,
Frames,
Brick,
Mesh,
Proxy,
Points,
Instances,
Attribute,
Relations,
}
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug, Serialize, Deserialize)]
pub struct PayloadReference {
pub kind: ReferencedPayloadKind,
pub dataset: u64,
pub chunk: u64,
pub byte_len: u64,
pub address: ContentAddress,
}
#[derive(Clone, PartialEq, Debug, Serialize, Deserialize)]
pub struct SceneManifest {
pub scene: SceneDescription,
pub payloads: Vec<PayloadReference>,
pub merkle_root: ContentAddress,
}
impl SceneManifest {
#[must_use]
pub fn new(scene: SceneDescription, mut payloads: Vec<PayloadReference>) -> Self {
payloads.sort_unstable();
let merkle_root = merkle_root(&payloads);
Self {
scene,
payloads,
merkle_root,
}
}
fn validate(&self) -> Result<(), ManifestError> {
if !self.payloads.windows(2).all(|pair| pair[0] <= pair[1]) {
return Err(ManifestError::NonCanonicalPayloadOrder);
}
let actual = merkle_root(&self.payloads);
if actual != self.merkle_root {
return Err(ManifestError::MerkleMismatch);
}
Ok(())
}
}
impl Scene {
#[must_use]
pub fn manifest(&self, mut payloads: Vec<PayloadReference>) -> SceneManifest {
let scene = self.describe();
payloads.extend(scene.point_batches.iter().map(|value| value.payload));
payloads.extend(scene.instance_batches.iter().map(|value| value.payload));
payloads.extend(scene.attributes.iter().map(|value| value.payload));
payloads.extend(scene.relation_batches.iter().map(|value| value.payload));
payloads.sort_unstable();
payloads.dedup();
SceneManifest::new(scene, payloads)
}
}
#[derive(Debug, Error)]
pub enum ManifestError {
#[error("manifest I/O failed: {0}")]
Io(#[from] io::Error),
#[error("manifest JSON failed: {0}")]
Json(#[from] serde_json::Error),
#[error("manifest payload references are not canonically ordered")]
NonCanonicalPayloadOrder,
#[error("manifest payload Merkle root does not match")]
MerkleMismatch,
#[error("payload is not available from the resolver")]
MissingPayload,
}
pub fn read_manifest<R: Read>(reader: R) -> Result<SceneManifest, ManifestError> {
let manifest: SceneManifest = serde_json::from_reader(reader)?;
manifest.validate()?;
Ok(manifest)
}
pub fn write_manifest<W: Write>(writer: W, manifest: &SceneManifest) -> Result<(), ManifestError> {
manifest.validate()?;
serde_json::to_writer(writer, manifest)?;
Ok(())
}
pub trait PayloadResolver {
type Reader: Read;
fn open(&self, reference: &PayloadReference) -> Result<Option<Self::Reader>, ManifestError>;
}
#[derive(Debug)]
pub struct LazyManifest<R> {
manifest: SceneManifest,
resolver: R,
}
impl<R: PayloadResolver> LazyManifest<R> {
pub fn new(manifest: SceneManifest, resolver: R) -> Result<Self, ManifestError> {
manifest.validate()?;
Ok(Self { manifest, resolver })
}
#[must_use]
pub const fn manifest(&self) -> &SceneManifest {
&self.manifest
}
pub fn open(
&self,
reference: &PayloadReference,
) -> Result<VerifiedPayload<R::Reader>, ManifestError> {
let Some(reader) = self.resolver.open(reference)? else {
return Err(ManifestError::MissingPayload);
};
Ok(VerifiedPayload::new(reader, *reference))
}
}
#[derive(Debug)]
pub struct VerifiedPayload<R> {
inner: R,
expected: PayloadReference,
hasher: Sha256,
bytes_read: u64,
verified: bool,
}
impl<R> VerifiedPayload<R> {
fn new(inner: R, expected: PayloadReference) -> Self {
Self {
inner,
expected,
hasher: Sha256::new(),
bytes_read: 0,
verified: false,
}
}
}
impl<R: Read> Read for VerifiedPayload<R> {
fn read(&mut self, buffer: &mut [u8]) -> io::Result<usize> {
let count = self.inner.read(buffer)?;
if count == 0 {
if !self.verified {
let actual: [u8; ADDRESS_BYTES] = self.hasher.clone().finalize().into();
if self.bytes_read != self.expected.byte_len || actual != self.expected.address.0 {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"resolved payload does not match its content address",
));
}
self.verified = true;
}
return Ok(0);
}
let increment = u64::try_from(count)
.map_err(|_| io::Error::other("payload read length exceeds u64"))?;
self.bytes_read = self
.bytes_read
.checked_add(increment)
.ok_or_else(|| io::Error::other("payload read length overflow"))?;
self.hasher.update(&buffer[..count]);
Ok(count)
}
}
fn merkle_root(payloads: &[PayloadReference]) -> ContentAddress {
let mut level: Vec<ContentAddress> = payloads.iter().map(reference_hash).collect();
if level.is_empty() {
return hash_parts(&[b"molgfx-empty-merkle"]);
}
while level.len() > 1 {
let mut parents = Vec::with_capacity(level.len().div_ceil(2));
for pair in level.chunks(2) {
let right = if pair.len() == 2 { pair[1] } else { pair[0] };
parents.push(hash_parts(&[b"molgfx-merkle-node", &pair[0].0, &right.0]));
}
level = parents;
}
level[0]
}
fn reference_hash(reference: &PayloadReference) -> ContentAddress {
let kind = [reference.kind as u8];
hash_parts(&[
b"molgfx-payload-reference",
&kind,
&reference.dataset.to_le_bytes(),
&reference.chunk.to_le_bytes(),
&reference.byte_len.to_le_bytes(),
&reference.address.0,
])
}
fn hash_parts(parts: &[&[u8]]) -> ContentAddress {
let mut hasher = Sha256::new();
for part in parts {
hasher.update(part);
}
ContentAddress(hasher.finalize().into())
}
#[cfg(test)]
#[path = "manifest_io_tests.rs"]
mod tests;