use crate::{Error, Result};
pub(crate) const TAG_INSERT: u8 = 0;
pub(crate) const TAG_REMOVE: u8 = 1;
pub(crate) const TAG_NAMESPACE_NAME: u8 = 2;
pub(crate) const TAG_ENCRYPTED_FLAG: u8 = 0x80;
pub(crate) const TAG_KIND_MASK: u8 = 0x7F;
pub(crate) const NONCE_LEN: usize = 12;
pub(crate) const TAG_LEN: usize = 16;
#[derive(Debug)]
pub(crate) enum RecordView<'a> {
Insert {
ns_id: u32,
key: &'a [u8],
value: &'a [u8],
expires_at: u64,
},
Remove {
ns_id: u32,
key: &'a [u8],
},
NamespaceName {
ns_id: u32,
name: &'a [u8],
},
}
#[derive(Debug)]
pub(crate) enum OwnedRecord {
Insert {
ns_id: u32,
key: Vec<u8>,
value: Vec<u8>,
expires_at: u64,
},
Remove {
ns_id: u32,
key: Vec<u8>,
},
NamespaceName {
ns_id: u32,
name: Vec<u8>,
},
}
#[inline]
pub(crate) fn write_u32(buf: &mut Vec<u8>, value: u32) {
buf.extend_from_slice(&value.to_le_bytes());
}
#[inline]
pub(crate) fn write_u64(buf: &mut Vec<u8>, value: u64) {
buf.extend_from_slice(&value.to_le_bytes());
}
#[inline]
fn checked_end(start: usize, len: usize, buf_len: usize) -> Option<usize> {
match start.checked_add(len) {
Some(end) if end <= buf_len => Some(end),
_ => None,
}
}
#[inline]
pub(crate) fn read_u32(bytes: &[u8], offset: usize) -> Result<u32> {
let end = checked_end(offset, 4, bytes.len()).ok_or(Error::Corrupted {
offset: offset as u64,
reason: "u32 read past end of buffer",
})?;
let mut buf = [0_u8; 4];
buf.copy_from_slice(&bytes[offset..end]);
Ok(u32::from_le_bytes(buf))
}
#[inline]
pub(crate) fn read_u64(bytes: &[u8], offset: usize) -> Result<u64> {
let end = checked_end(offset, 8, bytes.len()).ok_or(Error::Corrupted {
offset: offset as u64,
reason: "u64 read past end of buffer",
})?;
let mut buf = [0_u8; 8];
buf.copy_from_slice(&bytes[offset..end]);
Ok(u64::from_le_bytes(buf))
}
#[inline]
fn read_len_prefixed<'a>(
body: &'a [u8],
offset: usize,
reason: &'static str,
) -> Result<(&'a [u8], usize)> {
let len = read_u32(body, offset)? as usize;
let truncated = || Error::Corrupted {
offset: offset as u64,
reason,
};
let start = checked_end(offset, 4, body.len()).ok_or_else(truncated)?;
let end = checked_end(start, len, body.len()).ok_or_else(truncated)?;
Ok((&body[start..end], end))
}
#[inline]
fn require_exact_end(body: &[u8], end: usize, reason: &'static str) -> Result<()> {
if end == body.len() {
Ok(())
} else {
Err(Error::Corrupted {
offset: end as u64,
reason,
})
}
}
pub(crate) fn encode_insert_body(
out: &mut Vec<u8>,
ns_id: u32,
key: &[u8],
value: &[u8],
expires_at: u64,
) {
write_u32(out, ns_id);
write_u32(out, key.len() as u32);
out.extend_from_slice(key);
write_u32(out, value.len() as u32);
out.extend_from_slice(value);
write_u64(out, expires_at);
}
pub(crate) fn encode_remove_body(out: &mut Vec<u8>, ns_id: u32, key: &[u8]) {
write_u32(out, ns_id);
write_u32(out, key.len() as u32);
out.extend_from_slice(key);
}
pub(crate) fn encode_namespace_name_body(out: &mut Vec<u8>, ns_id: u32, name: &[u8]) {
write_u32(out, ns_id);
write_u32(out, name.len() as u32);
out.extend_from_slice(name);
}
pub(crate) fn decode_insert_body(body: &[u8]) -> Result<RecordView<'_>> {
let ns_id = read_u32(body, 0)?;
let (key, key_end) = read_len_prefixed(body, 4, "insert body truncated mid-key")?;
let (value, value_end) = read_len_prefixed(body, key_end, "insert body truncated mid-value")?;
let expires_at = read_u64(body, value_end)?;
let end = checked_end(value_end, 8, body.len()).ok_or(Error::Corrupted {
offset: value_end as u64,
reason: "u64 read past end of buffer",
})?;
require_exact_end(body, end, "insert body has trailing bytes")?;
Ok(RecordView::Insert {
ns_id,
key,
value,
expires_at,
})
}
pub(crate) fn decode_remove_body(body: &[u8]) -> Result<RecordView<'_>> {
let ns_id = read_u32(body, 0)?;
let (key, key_end) = read_len_prefixed(body, 4, "remove body truncated mid-key")?;
require_exact_end(body, key_end, "remove body has trailing bytes")?;
Ok(RecordView::Remove { ns_id, key })
}
pub(crate) fn decode_namespace_name_body(body: &[u8]) -> Result<RecordView<'_>> {
let ns_id = read_u32(body, 0)?;
let (name, name_end) = read_len_prefixed(body, 4, "namespace-name body truncated mid-name")?;
require_exact_end(body, name_end, "namespace-name body has trailing bytes")?;
Ok(RecordView::NamespaceName { ns_id, name })
}
pub(crate) fn decode_payload(payload: &[u8]) -> Result<RecordView<'_>> {
if payload.is_empty() {
return Err(Error::Corrupted {
offset: 0,
reason: "empty record payload",
});
}
let tag = payload[0];
if (tag & TAG_ENCRYPTED_FLAG) != 0 {
return Err(Error::Corrupted {
offset: 0,
reason: "encrypted record passed to plaintext decoder",
});
}
let body = &payload[1..];
match tag & TAG_KIND_MASK {
TAG_INSERT => decode_insert_body(body),
TAG_REMOVE => decode_remove_body(body),
TAG_NAMESPACE_NAME => decode_namespace_name_body(body),
unknown => Err(Error::Corrupted {
offset: 0,
reason: kind_error_for(unknown),
}),
}
}
pub(crate) fn decode_payload_encrypted<F>(payload: &[u8], decrypt: F) -> Result<OwnedRecord>
where
F: FnOnce(&[u8; NONCE_LEN], &[u8]) -> Result<Vec<u8>>,
{
if payload.len() < 1 + NONCE_LEN + TAG_LEN {
return Err(Error::Corrupted {
offset: 0,
reason: "encrypted payload shorter than nonce + AEAD tag",
});
}
let tag = payload[0];
if (tag & TAG_ENCRYPTED_FLAG) == 0 {
return Err(Error::Corrupted {
offset: 0,
reason: "plaintext record passed to encrypted decoder",
});
}
let kind = tag & TAG_KIND_MASK;
let mut nonce = [0_u8; NONCE_LEN];
nonce.copy_from_slice(&payload[1..1 + NONCE_LEN]);
let ciphertext = &payload[1 + NONCE_LEN..];
let plaintext = match decrypt(&nonce, ciphertext) {
Ok(p) => PlaintextBuf(p),
#[cfg(feature = "encrypt")]
Err(Error::EncryptionKeyMismatch) => {
return Err(Error::Corrupted {
offset: 0,
reason: "encrypted record failed authentication (modified or damaged)",
});
}
Err(err) => return Err(err),
};
let plaintext: &[u8] = &plaintext.0;
match kind {
TAG_INSERT => match decode_insert_body(plaintext)? {
RecordView::Insert {
ns_id,
key,
value,
expires_at,
} => Ok(OwnedRecord::Insert {
ns_id,
key: key.to_vec(),
value: value.to_vec(),
expires_at,
}),
_ => Err(Error::Corrupted {
offset: 0,
reason: "encrypted body shape mismatched its tag",
}),
},
TAG_REMOVE => match decode_remove_body(plaintext)? {
RecordView::Remove { ns_id, key } => Ok(OwnedRecord::Remove {
ns_id,
key: key.to_vec(),
}),
_ => Err(Error::Corrupted {
offset: 0,
reason: "encrypted body shape mismatched its tag",
}),
},
TAG_NAMESPACE_NAME => match decode_namespace_name_body(plaintext)? {
RecordView::NamespaceName { ns_id, name } => Ok(OwnedRecord::NamespaceName {
ns_id,
name: name.to_vec(),
}),
_ => Err(Error::Corrupted {
offset: 0,
reason: "encrypted body shape mismatched its tag",
}),
},
unknown => Err(Error::Corrupted {
offset: 0,
reason: kind_error_for(unknown),
}),
}
}
struct PlaintextBuf(Vec<u8>);
impl Drop for PlaintextBuf {
fn drop(&mut self) {
#[cfg(feature = "encrypt")]
zeroize::Zeroize::zeroize(&mut self.0);
}
}
pub(crate) fn payload_len_at(bytes: &[u8], payload_start: usize) -> Result<usize> {
if payload_start < 4 {
return Err(Error::Corrupted {
offset: payload_start as u64,
reason: "payload_start within frame header",
});
}
if payload_start > bytes.len() {
return Err(Error::Corrupted {
offset: payload_start as u64,
reason: "payload_start past buffer end",
});
}
Ok(read_u32(bytes, payload_start - 4)? as usize)
}
pub(crate) fn payload_at(bytes: &[u8], payload_start: usize) -> Result<&[u8]> {
let len = payload_len_at(bytes, payload_start)?;
let end = payload_start.checked_add(len).ok_or(Error::Corrupted {
offset: payload_start as u64,
reason: "payload_start + length overflowed",
})?;
if end > bytes.len() {
return Err(Error::Corrupted {
offset: payload_start as u64,
reason: "payload extends past buffer end",
});
}
Ok(&bytes[payload_start..end])
}
#[inline]
fn kind_error_for(_kind: u8) -> &'static str {
"unknown record tag kind"
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn insert_body_round_trips() {
let mut body = Vec::new();
encode_insert_body(&mut body, 7, b"key-bytes", b"value-bytes", 12345);
match decode_insert_body(&body).expect("decode") {
RecordView::Insert {
ns_id,
key,
value,
expires_at,
} => {
assert_eq!(ns_id, 7);
assert_eq!(key, b"key-bytes");
assert_eq!(value, b"value-bytes");
assert_eq!(expires_at, 12345);
}
_ => panic!("expected Insert"),
}
}
#[test]
fn payload_round_trips_via_decode_payload() {
let mut payload = vec![TAG_INSERT];
encode_insert_body(&mut payload, 0, b"k", b"v", 0);
match decode_payload(&payload).expect("decode") {
RecordView::Insert {
ns_id,
key,
value,
expires_at,
} => {
assert_eq!(ns_id, 0);
assert_eq!(key, b"k");
assert_eq!(value, b"v");
assert_eq!(expires_at, 0);
}
_ => panic!("expected Insert"),
}
}
#[test]
fn empty_payload_errors() {
let result = decode_payload(&[]);
assert!(matches!(result, Err(Error::Corrupted { .. })));
}
#[test]
fn unknown_tag_errors() {
let result = decode_payload(&[0x42_u8, 0, 0, 0, 0, 0, 0, 0, 0]);
assert!(matches!(result, Err(Error::Corrupted { .. })));
}
#[test]
fn encrypted_tag_to_plaintext_decoder_errors() {
let payload = vec![TAG_INSERT | TAG_ENCRYPTED_FLAG];
let result = decode_payload(&payload);
assert!(matches!(result, Err(Error::Corrupted { .. })));
}
fn insert_body(key: &[u8], value: &[u8]) -> Vec<u8> {
let mut body = Vec::new();
encode_insert_body(&mut body, 3, key, value, 99);
body
}
#[test]
fn test_decoders_accept_every_encoder_output_exactly() {
let body = insert_body(b"", b"");
assert!(decode_insert_body(&body).is_ok());
let body = insert_body(b"k", b"value");
assert!(decode_insert_body(&body).is_ok());
let mut rm = Vec::new();
encode_remove_body(&mut rm, 1, b"k");
assert!(decode_remove_body(&rm).is_ok());
let mut empty_rm = Vec::new();
encode_remove_body(&mut empty_rm, 0, b"");
assert!(decode_remove_body(&empty_rm).is_ok());
let mut nn = Vec::new();
encode_namespace_name_body(&mut nn, 1, b"users");
assert!(decode_namespace_name_body(&nn).is_ok());
}
#[test]
fn test_decoders_with_trailing_bytes_return_corrupted() {
let mut body = insert_body(b"k", b"v");
body.push(0);
assert!(matches!(
decode_insert_body(&body),
Err(Error::Corrupted { reason, .. }) if reason.contains("trailing")
));
let mut rm = Vec::new();
encode_remove_body(&mut rm, 1, b"k");
rm.push(0);
assert!(matches!(
decode_remove_body(&rm),
Err(Error::Corrupted { .. })
));
let mut nn = Vec::new();
encode_namespace_name_body(&mut nn, 1, b"users");
nn.push(0);
assert!(matches!(
decode_namespace_name_body(&nn),
Err(Error::Corrupted { .. })
));
}
#[test]
fn test_insert_body_decoded_as_other_kind_returns_corrupted() {
for (key, value) in [(&b"victim"[..], &b"important-value"[..]), (b"", b"")] {
let body = insert_body(key, value);
assert!(decode_remove_body(&body).is_err());
assert!(decode_namespace_name_body(&body).is_err());
}
let mut rm = Vec::new();
encode_remove_body(&mut rm, 1, b"some-key");
assert!(decode_insert_body(&rm).is_err());
}
#[test]
fn test_decoders_with_max_lengths_do_not_overflow() {
let mut body = Vec::new();
write_u32(&mut body, 0);
write_u32(&mut body, u32::MAX);
body.extend_from_slice(&[0_u8; 16]);
assert!(matches!(
decode_insert_body(&body),
Err(Error::Corrupted { .. })
));
assert!(matches!(
decode_remove_body(&body),
Err(Error::Corrupted { .. })
));
assert!(matches!(
decode_namespace_name_body(&body),
Err(Error::Corrupted { .. })
));
let mut body = Vec::new();
write_u32(&mut body, 0);
write_u32(&mut body, 1);
body.push(b'k');
write_u32(&mut body, u32::MAX);
assert!(matches!(
decode_insert_body(&body),
Err(Error::Corrupted { .. })
));
assert!(read_u32(&[0_u8; 4], usize::MAX).is_err());
assert!(read_u64(&[0_u8; 8], usize::MAX - 3).is_err());
assert!(payload_at(&[0_u8; 8], usize::MAX).is_err());
}
#[test]
fn test_decode_payload_encrypted_rejects_plaintext_tag() {
let mut payload = vec![TAG_REMOVE];
payload.extend_from_slice(&[0_u8; NONCE_LEN + TAG_LEN + 8]);
let result = decode_payload_encrypted(&payload, |_, _| Ok(Vec::new()));
assert!(matches!(result, Err(Error::Corrupted { .. })));
}
#[cfg(feature = "encrypt")]
#[test]
fn test_decode_payload_encrypted_auth_failure_returns_corrupted() {
let mut payload = vec![TAG_INSERT | TAG_ENCRYPTED_FLAG];
payload.extend_from_slice(&[0_u8; NONCE_LEN + TAG_LEN + 8]);
let result = decode_payload_encrypted(&payload, |_, _| Err(Error::EncryptionKeyMismatch));
assert!(matches!(result, Err(Error::Corrupted { .. })), "{result:?}");
}
#[test]
fn payload_at_handles_basic_geometry() {
let mut frame = Vec::new();
frame.extend_from_slice(&0x4653_5901_u32.to_be_bytes()); frame.extend_from_slice(&5_u32.to_le_bytes()); frame.extend_from_slice(b"hello"); frame.extend_from_slice(&0_u32.to_le_bytes());
let payload_start = 8;
let payload = payload_at(&frame, payload_start).expect("payload_at");
assert_eq!(payload, b"hello");
assert_eq!(payload_len_at(&frame, payload_start).expect("len"), 5);
}
}