use crate::digest::sha256_digest;
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use serde_json::value::RawValue;
use thiserror::Error;
#[derive(Debug, Deserialize)]
pub struct EnvelopeProbe {
pub kind: String,
pub format_version: u32,
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum EnvelopeCodecError {
#[error("failed to encode envelope payload: {0}")]
PayloadEncode(String),
#[error("failed to encode envelope document: {0}")]
EnvelopeEncode(String),
#[error("failed to decode envelope document: {0}")]
EnvelopeDecode(String),
#[error("failed to decode envelope payload: {0}")]
PayloadDecode(String),
#[error("failed to compress envelope: {0}")]
Compress(String),
#[error("failed to decompress envelope: {0}")]
Decompress(String),
#[error("decompressed WAL document exceeds the {max_bytes}-byte limit")]
WalSegmentTooLarge {
max_bytes: usize,
},
#[error(
"invalid wal inline content in commit `{seq}` for `content_id` `{content_id}`: {reason}"
)]
InvalidWalInlineContent {
seq: crate::ChangeSeq,
content_id: crate::ContentId,
reason: &'static str,
},
#[error("unknown envelope kind `{found}`")]
UnknownKind {
found: String,
},
#[error("envelope kind mismatch: expected `{expected}`, found `{found}`")]
KindMismatch {
expected: String,
found: String,
},
#[error(
"unsupported `{kind}` envelope format version `{found}`: \
this build supports `{supported}`"
)]
UnsupportedFormatVersion {
kind: String,
found: u32,
supported: u32,
},
#[error("envelope payload checksum mismatch: expected `{expected}`, actual `{actual}`")]
ChecksumMismatch {
expected: String,
actual: String,
},
}
pub fn verify_kind(expected: &str, found: &str) -> Result<(), EnvelopeCodecError> {
if found != expected {
return Err(EnvelopeCodecError::KindMismatch {
expected: expected.to_owned(),
found: found.to_owned(),
});
}
Ok(())
}
pub fn verify_version(kind: &str, found: u32, supported: u32) -> Result<(), EnvelopeCodecError> {
if found != supported {
return Err(EnvelopeCodecError::UnsupportedFormatVersion {
kind: kind.to_owned(),
found,
supported,
});
}
Ok(())
}
pub fn verify_payload_checksum(
expected: &str,
payload_bytes: &[u8],
) -> Result<(), EnvelopeCodecError> {
let actual = sha256_digest(payload_bytes);
if actual != expected {
return Err(EnvelopeCodecError::ChecksumMismatch {
expected: expected.to_owned(),
actual,
});
}
Ok(())
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct JsonEnvelopeDocument {
kind: String,
format_version: u32,
payload_checksum: String,
payload: Box<RawValue>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VerifiedEnvelope<T> {
pub(crate) payload_checksum: String,
pub(crate) payload: T,
}
impl<T> VerifiedEnvelope<T> {
pub fn payload_checksum(&self) -> &str {
&self.payload_checksum
}
pub fn payload(&self) -> &T {
&self.payload
}
pub fn into_payload(self) -> T {
self.payload
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EncodedEnvelope<T> {
pub(crate) envelope: VerifiedEnvelope<T>,
pub(crate) bytes: Vec<u8>,
pub(crate) document_len: usize,
}
impl<T> EncodedEnvelope<T> {
pub fn document_len(&self) -> usize {
self.document_len
}
pub fn envelope(&self) -> &VerifiedEnvelope<T> {
&self.envelope
}
pub fn as_bytes(&self) -> &[u8] {
&self.bytes
}
pub fn into_bytes(self) -> Vec<u8> {
self.bytes
}
pub fn into_envelope(self) -> VerifiedEnvelope<T> {
self.envelope
}
pub fn into_parts(self) -> (VerifiedEnvelope<T>, Vec<u8>) {
(self.envelope, self.bytes)
}
}
pub fn encode_json_envelope<T: Serialize>(
kind: &str,
format_version: u32,
payload: T,
) -> Result<EncodedEnvelope<T>, EnvelopeCodecError> {
let payload_json = serde_json::to_string(&payload)
.map_err(|err| EnvelopeCodecError::PayloadEncode(err.to_string()))?;
let payload_checksum = sha256_digest(payload_json.as_bytes());
let document = JsonEnvelopeDocument {
kind: kind.to_owned(),
format_version,
payload_checksum: payload_checksum.clone(),
payload: RawValue::from_string(payload_json)
.map_err(|err| EnvelopeCodecError::PayloadEncode(err.to_string()))?,
};
let bytes = serde_json::to_vec(&document)
.map_err(|err| EnvelopeCodecError::EnvelopeEncode(err.to_string()))?;
Ok(EncodedEnvelope {
envelope: VerifiedEnvelope {
payload_checksum,
payload,
},
document_len: bytes.len(),
bytes,
})
}
pub fn decode_json_envelope<T: DeserializeOwned>(
bytes: &[u8],
supported_version: u32,
classify_kind: impl FnOnce(&str) -> Result<(), EnvelopeCodecError>,
) -> Result<VerifiedEnvelope<T>, EnvelopeCodecError> {
let probe: EnvelopeProbe = serde_json::from_slice(bytes)
.map_err(|err| EnvelopeCodecError::EnvelopeDecode(err.to_string()))?;
classify_kind(&probe.kind)?;
verify_version(&probe.kind, probe.format_version, supported_version)?;
let document: JsonEnvelopeDocument = serde_json::from_slice(bytes)
.map_err(|err| EnvelopeCodecError::EnvelopeDecode(err.to_string()))?;
decode_json_envelope_payload(document)
}
fn decode_json_envelope_payload<T: DeserializeOwned>(
document: JsonEnvelopeDocument,
) -> Result<VerifiedEnvelope<T>, EnvelopeCodecError> {
verify_payload_checksum(
&document.payload_checksum,
document.payload.get().as_bytes(),
)?;
let payload: T = serde_json::from_str(document.payload.get())
.map_err(|err| EnvelopeCodecError::PayloadDecode(err.to_string()))?;
Ok(VerifiedEnvelope {
payload_checksum: document.payload_checksum,
payload,
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::Cell;
#[test]
fn encoding_serializes_the_payload_once_and_checksums_those_bytes() {
struct CountedPayload<'a>(&'a Cell<usize>);
impl Serialize for CountedPayload<'_> {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.0.set(self.0.get() + 1);
serializer.serialize_u64(42)
}
}
let calls = Cell::new(0);
let encoded = encode_json_envelope("test", 1, CountedPayload(&calls)).expect("encode");
let document: JsonEnvelopeDocument =
serde_json::from_slice(encoded.as_bytes()).expect("document");
assert_eq!(calls.get(), 1);
assert_eq!(document.payload.get(), "42");
assert_eq!(encoded.envelope().payload_checksum(), sha256_digest(b"42"));
assert_eq!(
document.payload_checksum,
encoded.envelope().payload_checksum()
);
}
#[test]
fn decoding_checks_the_stored_payload_including_noncanonical_whitespace() {
let payload = r#"{ "value" : 42 }"#;
let checksum = sha256_digest(payload.as_bytes());
let bytes = format!(
r#"{{"kind":"test","format_version":1,"payload_checksum":"{checksum}","payload":{payload}}}"#
);
let decoded: VerifiedEnvelope<serde_json::Value> =
decode_json_envelope(bytes.as_bytes(), 1, |kind| verify_kind("test", kind))
.expect("decode exact stored bytes");
assert_eq!(decoded.payload_checksum(), checksum);
let successor =
encode_json_envelope("test", 1, decoded.into_payload()).expect("canonical encoding");
assert_ne!(successor.envelope().payload_checksum(), checksum);
let reread: VerifiedEnvelope<serde_json::Value> =
decode_json_envelope(successor.as_bytes(), 1, |kind| verify_kind("test", kind))
.expect("decode canonical bytes");
assert_eq!(&reread, successor.envelope());
}
}