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)]
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("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,
},
#[error(
"envelope checksum `{checksum}` does not match its payload `{actual}`: \
rebuild the envelope from its payload"
)]
StalePayloadChecksum {
checksum: 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(())
}
pub fn verify_checksum_fresh(
checksum: &str,
payload_bytes: &[u8],
) -> Result<(), EnvelopeCodecError> {
let actual = sha256_digest(payload_bytes);
if actual != checksum {
return Err(EnvelopeCodecError::StalePayloadChecksum {
checksum: checksum.to_owned(),
actual,
});
}
Ok(())
}
#[derive(Serialize, Deserialize)]
struct JsonEnvelopeDocument {
kind: String,
format_version: u32,
payload_checksum: String,
payload: Box<RawValue>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct StrictJsonEnvelopeDocument {
kind: String,
format_version: u32,
payload_checksum: String,
payload: Box<RawValue>,
}
impl From<StrictJsonEnvelopeDocument> for JsonEnvelopeDocument {
fn from(document: StrictJsonEnvelopeDocument) -> Self {
Self {
kind: document.kind,
format_version: document.format_version,
payload_checksum: document.payload_checksum,
payload: document.payload,
}
}
}
pub fn json_payload_checksum<T: Serialize>(payload: &T) -> Result<String, EnvelopeCodecError> {
let bytes = serde_json::to_vec(payload)
.map_err(|err| EnvelopeCodecError::PayloadEncode(err.to_string()))?;
Ok(sha256_digest(&bytes))
}
pub fn encode_json_envelope<T: Serialize>(
kind: &str,
format_version: u32,
supported_version: u32,
payload_checksum: &str,
payload: &T,
) -> Result<Vec<u8>, EnvelopeCodecError> {
verify_version(kind, format_version, supported_version)?;
let payload_json = serde_json::to_string(payload)
.map_err(|err| EnvelopeCodecError::PayloadEncode(err.to_string()))?;
verify_checksum_fresh(payload_checksum, payload_json.as_bytes())?;
let document = JsonEnvelopeDocument {
kind: kind.to_owned(),
format_version,
payload_checksum: payload_checksum.to_owned(),
payload: RawValue::from_string(payload_json)
.map_err(|err| EnvelopeCodecError::PayloadEncode(err.to_string()))?,
};
serde_json::to_vec(&document).map_err(|err| EnvelopeCodecError::EnvelopeEncode(err.to_string()))
}
pub struct DecodedJsonEnvelope<T> {
pub format_version: u32,
pub payload_checksum: String,
pub payload: T,
}
pub fn decode_json_envelope<T: DeserializeOwned>(
bytes: &[u8],
supported_version: u32,
classify_kind: impl FnOnce(&str) -> Result<(), EnvelopeCodecError>,
) -> Result<DecodedJsonEnvelope<T>, EnvelopeCodecError> {
decode_json_envelope_probe(bytes, supported_version, classify_kind)?;
let document: JsonEnvelopeDocument = serde_json::from_slice(bytes)
.map_err(|err| EnvelopeCodecError::EnvelopeDecode(err.to_string()))?;
decode_json_envelope_payload(document)
}
pub fn decode_strict_json_envelope<T: DeserializeOwned>(
bytes: &[u8],
supported_version: u32,
classify_kind: impl FnOnce(&str) -> Result<(), EnvelopeCodecError>,
) -> Result<DecodedJsonEnvelope<T>, EnvelopeCodecError> {
decode_json_envelope_probe(bytes, supported_version, classify_kind)?;
let document: StrictJsonEnvelopeDocument = serde_json::from_slice(bytes)
.map_err(|err| EnvelopeCodecError::EnvelopeDecode(err.to_string()))?;
decode_json_envelope_payload(document.into())
}
fn decode_json_envelope_probe(
bytes: &[u8],
supported_version: u32,
classify_kind: impl FnOnce(&str) -> Result<(), EnvelopeCodecError>,
) -> Result<(), 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)?;
Ok(())
}
fn decode_json_envelope_payload<T: DeserializeOwned>(
document: JsonEnvelopeDocument,
) -> Result<DecodedJsonEnvelope<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(DecodedJsonEnvelope {
format_version: document.format_version,
payload_checksum: document.payload_checksum,
payload,
})
}