use core::result;
pub const AOF_HEADER_LEN: usize = 8;
const KEY_LEN_PREFIX: usize = 4;
const BLOB_LEN_PREFIX: usize = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum AofOp {
KvUpsert = 1,
KvDelete = 2,
RiCreate = 3,
RiSet = 4,
RiDel = 5,
TtlPurge = 6,
}
impl TryFrom<u8> for AofOp {
type Error = Error;
#[inline]
fn try_from(v: u8) -> AofResult<Self> {
match v {
1 => Ok(Self::KvUpsert),
2 => Ok(Self::KvDelete),
3 => Ok(Self::RiCreate),
4 => Ok(Self::RiSet),
5 => Ok(Self::RiDel),
6 => Ok(Self::TtlPurge),
_ => Err(Error::UnknownOp(v)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AofEntryRef<'a> {
pub op: AofOp,
pub version: u32,
pub key: &'a [u8],
pub blob: &'a [u8],
}
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("AOF entry truncated: need {need} bytes, got {got}")]
Truncated { need: usize, got: usize },
#[error("AOF entry length prefix overflow: declared {declared}, remaining {remaining}")]
Overflow { declared: usize, remaining: usize },
#[error("unknown AOF op: {0}")]
UnknownOp(u8),
}
pub type AofResult<T> = result::Result<T, Error>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TtlPurgePayload {
pub ns: u64,
pub db: u64,
pub expire_at_ms: u64,
}
impl TtlPurgePayload {
pub const LEN: usize = 24;
#[inline]
pub fn encode(self) -> [u8; Self::LEN] {
let mut buf = [0u8; Self::LEN];
buf[0..8].copy_from_slice(&self.ns.to_be_bytes());
buf[8..16].copy_from_slice(&self.db.to_be_bytes());
buf[16..24].copy_from_slice(&self.expire_at_ms.to_be_bytes());
buf
}
#[inline]
pub fn decode(blob: &[u8]) -> AofResult<Self> {
match blob.len() {
Self::LEN => {}
n if n < Self::LEN => {
return Err(Error::Truncated {
need: Self::LEN,
got: n,
});
}
n => {
return Err(Error::Overflow {
declared: Self::LEN,
remaining: n,
});
}
}
Ok(Self {
ns: u64::from_be_bytes(unsafe { blob.get_unchecked(0..8).try_into().unwrap_unchecked() }),
db: u64::from_be_bytes(unsafe { blob.get_unchecked(8..16).try_into().unwrap_unchecked() }),
expire_at_ms: u64::from_be_bytes(unsafe {
blob.get_unchecked(16..24).try_into().unwrap_unchecked()
}),
})
}
}
pub fn encode_entry(op: AofOp, version: u32, key: &[u8], blob: &[u8]) -> Vec<u8> {
let total = AOF_HEADER_LEN + KEY_LEN_PREFIX + key.len() + BLOB_LEN_PREFIX + blob.len();
let mut buf = Vec::with_capacity(total);
buf.push(op as u8);
buf.push(0); buf.extend_from_slice(&[0, 0]); buf.extend_from_slice(&version.to_le_bytes());
buf.extend_from_slice(&(key.len() as u32).to_le_bytes());
buf.extend_from_slice(key);
buf.extend_from_slice(&(blob.len() as u32).to_le_bytes());
buf.extend_from_slice(blob);
buf
}
impl<'a> AofEntryRef<'a> {
pub fn decode(buf: &'a [u8]) -> AofResult<Self> {
if buf.len() < AOF_HEADER_LEN {
return Err(Error::Truncated {
need: AOF_HEADER_LEN,
got: buf.len(),
});
}
let op = AofOp::try_from(buf[0])?;
let version =
u32::from_le_bytes(unsafe { buf.get_unchecked(4..8).try_into().unwrap_unchecked() });
let key_len = read_len_prefix(buf, AOF_HEADER_LEN, KEY_LEN_PREFIX)?;
let key_end = AOF_HEADER_LEN + KEY_LEN_PREFIX + key_len;
if buf.len() < key_end {
return Err(Error::Overflow {
declared: key_len,
remaining: buf.len() - AOF_HEADER_LEN - KEY_LEN_PREFIX,
});
}
let key = &buf[AOF_HEADER_LEN + KEY_LEN_PREFIX..key_end];
let blob_len = read_len_prefix(buf, key_end, BLOB_LEN_PREFIX)?;
let blob_end = key_end + BLOB_LEN_PREFIX + blob_len;
if buf.len() < blob_end {
return Err(Error::Overflow {
declared: blob_len,
remaining: buf.len() - key_end - BLOB_LEN_PREFIX,
});
}
let blob = &buf[key_end + BLOB_LEN_PREFIX..blob_end];
Ok(Self {
op,
version,
key,
blob,
})
}
}
#[inline]
fn read_len_prefix(buf: &[u8], offset: usize, width: usize) -> AofResult<usize> {
let end = offset + width;
if buf.len() < end {
return Err(Error::Truncated {
need: end,
got: buf.len(),
});
}
Ok(
u32::from_le_bytes(unsafe { buf.get_unchecked(offset..end).try_into().unwrap_unchecked() })
as usize,
)
}
pub trait Replay {
fn on_entry(&mut self, entry: AofEntryRef<'_>) -> AofResult<()>;
fn on_ttl_purge(&mut self, _key: &[u8], _payload: TtlPurgePayload) -> AofResult<()> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn encode_decode_round_trip() {
let buf = encode_entry(AofOp::RiSet, 7, b"idx", b"*4\r\n...");
assert_eq!(buf.len(), 8 + 4 + 3 + 4 + 7);
let e = AofEntryRef::decode(&buf).unwrap();
assert_eq!(e.op, AofOp::RiSet);
assert_eq!(e.version, 7);
assert_eq!(e.key, b"idx");
assert_eq!(e.blob, b"*4\r\n...");
}
#[test]
fn decode_empty_blob() {
let buf = encode_entry(AofOp::KvDelete, 0, b"k", &[]);
let e = AofEntryRef::decode(&buf).unwrap();
assert_eq!(e.op, AofOp::KvDelete);
assert_eq!(e.blob, &[] as &[u8]);
}
#[test]
fn decode_truncated_rejected() {
let buf = encode_entry(AofOp::RiDel, 1, b"idx", b"x");
for cut in [0, 4, 8, 8 + 2, 8 + 4 + 2] {
assert!(AofEntryRef::decode(&buf[..cut]).is_err(), "cut={cut}");
}
}
#[test]
fn decode_unknown_op_rejected() {
let mut buf = encode_entry(AofOp::RiDel, 1, b"idx", b"x");
buf[0] = 0xFF;
assert!(matches!(
AofEntryRef::decode(&buf),
Err(Error::UnknownOp(0xFF))
));
}
#[test]
fn ttl_purge_payload_round_trip() {
let payload = TtlPurgePayload {
ns: 7,
db: 3,
expire_at_ms: 0x0102_0304_0506_0708,
};
let blob = payload.encode();
assert_eq!(blob.len(), TtlPurgePayload::LEN);
assert_eq!(&blob[0..8], &7u64.to_be_bytes());
assert_eq!(&blob[8..16], &3u64.to_be_bytes());
assert_eq!(&blob[16..24], &0x0102_0304_0506_0708u64.to_be_bytes());
assert_eq!(TtlPurgePayload::decode(&blob).unwrap(), payload);
}
#[test]
fn ttl_purge_entry_round_trip_and_malformed_blob() {
let blob = TtlPurgePayload {
ns: 1,
db: 2,
expire_at_ms: 42,
}
.encode();
let buf = encode_entry(AofOp::TtlPurge, 9, b"glitch", &blob);
let e = AofEntryRef::decode(&buf).unwrap();
assert_eq!(e.op, AofOp::TtlPurge);
assert_eq!(e.key, b"glitch");
assert_eq!(TtlPurgePayload::decode(e.blob).unwrap().expire_at_ms, 42);
assert!(matches!(
TtlPurgePayload::decode(&blob[..23]),
Err(Error::Truncated { need: 24, got: 23 })
));
let mut long = blob.to_vec();
long.push(0);
assert!(matches!(
TtlPurgePayload::decode(&long),
Err(Error::Overflow {
declared: 24,
remaining: 25
})
));
}
}