use crate::error::{ResolveError, Result};
use crate::urn::ParsedUrn;
use base64::Engine;
use digstore_core::codec::Decode;
use digstore_core::crypto::{decrypt_chunk, derive_decryption_key};
use digstore_core::{resource_leaf, Bytes32, MerkleProof, SecretSalt};
fn parse_salt(salt_hex: Option<&str>) -> Result<Option<[u8; 32]>> {
match salt_hex {
None => Ok(None),
Some(s) if s.trim().is_empty() => Ok(None),
Some(s) => {
let b = Bytes32::from_hex(s.trim())
.map_err(|_| ResolveError::Parse("secret salt must be 64 hex chars".into()))?;
Ok(Some(b.0))
}
}
}
fn decode_proof_b64(proof_b64: &str) -> Result<MerkleProof> {
let raw = base64::engine::general_purpose::STANDARD
.decode(proof_b64.trim().as_bytes())
.map_err(|_| ResolveError::VerifyFailed("inclusion proof is not valid base64".into()))?;
MerkleProof::from_bytes(&raw)
.map_err(|_| ResolveError::VerifyFailed("inclusion proof encoding is invalid".into()))
}
pub fn verify_inclusion(ciphertext: &[u8], proof_b64: &str, trusted_root_hex: &str) -> Result<()> {
let trusted_root = Bytes32::from_hex(trusted_root_hex.trim())
.map_err(|_| ResolveError::VerifyFailed("trusted root must be 64 hex chars".into()))?;
let proof = decode_proof_b64(proof_b64)?;
if resource_leaf(ciphertext) != proof.leaf {
return Err(ResolveError::VerifyFailed(
"content does not match proof leaf (tampered ciphertext)".into(),
));
}
if !proof.verify() {
return Err(ResolveError::VerifyFailed(
"merkle path does not resolve to the declared root".into(),
));
}
if proof.root != trusted_root {
return Err(ResolveError::VerifyFailed(
"merkle root does not match the chain-anchored trusted root".into(),
));
}
Ok(())
}
pub fn decrypt(parsed: &ParsedUrn, ciphertext: &[u8], chunk_lens: &[u32]) -> Result<Vec<u8>> {
let salt = parse_salt(parsed.salt.as_deref())?;
let canonical = parsed.canonical_rootless().canonical();
let aes_key = derive_decryption_key(&canonical, salt.map(SecretSalt).as_ref());
let ct_len = ciphertext.len() as u64;
let plan: Vec<u64> = if chunk_lens.is_empty() {
vec![ct_len]
} else {
chunk_lens.iter().map(|&l| l as u64).collect()
};
let mut total: u64 = 0;
for &len in &plan {
total = total.checked_add(len).ok_or(ResolveError::DecryptFailed)?;
}
if total != ct_len {
return Err(ResolveError::DecryptFailed);
}
let mut plaintext = Vec::with_capacity(ciphertext.len());
let mut p: usize = 0;
for len in plan {
let len = usize::try_from(len).map_err(|_| ResolveError::DecryptFailed)?;
let end = p.checked_add(len).filter(|&e| e <= ciphertext.len());
let end = end.ok_or(ResolveError::DecryptFailed)?;
let ct = &ciphertext[p..end];
p = end;
let pt = decrypt_chunk(&aes_key, ct).map_err(|_| ResolveError::DecryptFailed)?;
plaintext.extend_from_slice(&pt);
}
Ok(plaintext)
}
pub fn verify_and_decrypt(
parsed: &ParsedUrn,
ciphertext: &[u8],
proof_b64: &str,
trusted_root_hex: &str,
chunk_lens: &[u32],
) -> Result<Vec<u8>> {
verify_inclusion(ciphertext, proof_b64, trusted_root_hex)?;
decrypt(parsed, ciphertext, chunk_lens)
}
#[cfg(test)]
mod tests {
use super::*;
fn urn() -> ParsedUrn {
ParsedUrn::parse(&format!("urn:dig:chia:{}/a.bin", "ab".repeat(32))).unwrap()
}
#[test]
fn decrypt_rejects_overflowing_chunk_lens_without_panic() {
let ct = vec![0u8; 10];
let bad = [(1u32 << 31) + 10, 1u32 << 31];
assert!(matches!(
decrypt(&urn(), &ct, &bad),
Err(ResolveError::DecryptFailed)
));
}
#[test]
fn decrypt_rejects_chunk_total_mismatch() {
let ct = vec![0u8; 10];
assert!(matches!(
decrypt(&urn(), &ct, &[999]),
Err(ResolveError::DecryptFailed)
));
}
}