#![deny(unsafe_code)]
pub mod arena;
pub mod client;
pub mod config;
pub mod constants;
pub mod envelope;
pub mod http_transport;
pub mod lifecycle;
pub mod runtime_tpc;
pub mod tcp_transport;
pub use arena::{
arena_path, ArenaError, ArenaReader, ArenaWriter, ARENA_HEADER_BYTES, DEFAULT_ARENA_BYTES,
};
pub use client::{
ClientOptions, RemoteError, RemoteGetOutcome, RemoteHitTier, RemoteKvStoreClient,
DEFAULT_CALL_TIMEOUT,
};
pub use config::DaemonConfig;
pub use lifecycle::{cleanup_prefix_segments, ClientHeartbeat, HeartbeatMonitor, ReopenReason};
use std::time::Duration;
use bytes::Bytes;
use myelon::codec::{Codec, CodecError};
use myelon::transport::AlignedFixedFrame;
use myelon::typed_transport::{TypedConsumer, TypedProducer};
use myelon::MyelonWaitStrategy;
use rkyv::rancor::Error as RkyvError;
use rkyv::util::AlignedVec;
use rkyv::{Archive, Deserialize as RkyvDeserialize, Serialize as RkyvSerialize};
pub const FRAME_DATA_BYTES: usize = 4 * 1024 * 1024 - 16;
pub const DEFAULT_RING_DEPTH: usize = 16;
pub type ShmFrame = AlignedFixedFrame<FRAME_DATA_BYTES>;
pub const REQ_CONSUMER_ID: &str = "dn";
pub const RESP_CONSUMER_ID: &str = "cn";
pub const ATTACH_TIMEOUT: Duration = Duration::from_secs(30);
#[must_use]
pub fn effective_attach_timeout() -> Duration {
std::env::var("WMBT_KV_DAEMON_SHM_ATTACH_TIMEOUT_SECS")
.ok()
.and_then(|s| s.parse::<u64>().ok())
.filter(|n| *n >= 1)
.map_or(ATTACH_TIMEOUT, Duration::from_secs)
}
pub mod op {
pub const PING: u8 = 1;
pub const PUT: u8 = 2;
pub const GET: u8 = 3;
pub const EXISTS: u8 = 4;
pub const STATS: u8 = 5;
pub const CLEAR: u8 = 6;
pub const RESTORE: u8 = 7;
pub const CLOSE: u8 = 8;
pub const GET_MANY: u8 = 9;
pub const LIST: u8 = 10;
pub const LOOKUP_BLOCK_PREFIX: u8 = 11;
pub const GET_KV_BLOCKS_BATCH: u8 = 12;
pub const PUT_KV_BLOCKS_BATCH: u8 = 13;
#[must_use]
pub fn name(code: u8) -> &'static str {
match code {
PING => "ping",
PUT => "put",
GET => "get",
EXISTS => "exists",
STATS => "stats",
CLEAR => "clear",
RESTORE => "restore",
CLOSE => "close",
GET_MANY => "get_many",
LIST => "list",
LOOKUP_BLOCK_PREFIX => "lookup_block_prefix",
GET_KV_BLOCKS_BATCH => "get_kv_blocks_batch",
PUT_KV_BLOCKS_BATCH => "put_kv_blocks_batch",
_ => "unknown",
}
}
}
pub mod status {
pub const OK: u8 = 0;
pub const MISS: u8 = 1;
pub const TOO_LARGE: u8 = 2;
pub const ERROR: u8 = 3;
}
const KEY_BATCH_MAGIC: &[u8] = b"WMBT_KV_KEYS_V1\0";
const BYTES_BATCH_MAGIC: &[u8] = b"WMBT_KV_BYTES_V1\0";
const LOOKUP_BLOCK_PREFIX_REQ_MAGIC: &[u8] = b"WMBT_LBP_REQ\0\0\0\0";
const LOOKUP_BLOCK_PREFIX_RESP_MAGIC: &[u8] = b"WMBT_LBP_RES\0\0\0\0";
const GET_KV_BLOCKS_BATCH_REQ_MAGIC: &[u8] = b"WMBT_GBB_REQ\0\0\0\0";
const GET_KV_BLOCKS_BATCH_RESP_MAGIC: &[u8] = b"WMBT_GBB_RES\0\0\0\0";
const PUT_KV_BLOCKS_BATCH_REQ_MAGIC: &[u8] = b"WMBT_PBB_REQ\0\0\0\0";
const PUT_KV_BLOCKS_BATCH_RESP_MAGIC: &[u8] = b"WMBT_PBB_RES\0\0\0\0";
#[must_use]
pub fn encode_key_batch(keys: &[String]) -> Vec<u8> {
let total = KEY_BATCH_MAGIC.len() + 4 + keys.iter().map(|key| 4 + key.len()).sum::<usize>();
let mut out = Vec::with_capacity(total);
out.extend_from_slice(KEY_BATCH_MAGIC);
out.extend_from_slice(&(keys.len() as u32).to_le_bytes());
for key in keys {
let bytes = key.as_bytes();
out.extend_from_slice(&(bytes.len() as u32).to_le_bytes());
out.extend_from_slice(bytes);
}
out
}
#[derive(Debug, thiserror::Error)]
pub enum WireCodecError {
#[error("{what} truncated: need {needed} bytes at offset {at}, payload is {got}")]
Truncated { what: &'static str, needed: usize, at: usize, got: usize },
#[error("{op}: bad magic; expected {expected:?} got first {got_len} bytes {got:?}")]
BadMagic { op: &'static str, expected: &'static [u8], got_len: usize, got: Vec<u8> },
#[error("{what}: length overflow")]
LengthOverflow { what: &'static str },
#[error("{what}: trailing bytes ({extra} unread)")]
TrailingBytes { what: &'static str, extra: usize },
#[error("{what}: body length mismatch (claimed={claimed}, actual={actual})")]
BodyLengthMismatch { what: &'static str, claimed: usize, actual: usize },
#[error("{what}: utf8: {source}")]
Utf8 {
what: &'static str,
#[source]
source: std::str::Utf8Error,
},
#[error("{op}: rkyv: {source}")]
Rkyv {
op: &'static str,
#[source]
source: rkyv::rancor::Error,
},
#[error("{0}")]
SegmentNameBudget(String),
}
pub fn decode_key_batch(payload: &[u8]) -> Result<Vec<String>, WireCodecError> {
if !payload.starts_with(KEY_BATCH_MAGIC) {
let got_len = KEY_BATCH_MAGIC.len().min(payload.len());
return Err(WireCodecError::BadMagic {
op: "key_batch",
expected: KEY_BATCH_MAGIC,
got_len,
got: payload[..got_len].to_vec(),
});
}
let mut cursor = KEY_BATCH_MAGIC.len();
let count = read_u32(payload, &mut cursor)? as usize;
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
let len = read_u32(payload, &mut cursor)? as usize;
let end =
cursor.checked_add(len).ok_or(WireCodecError::LengthOverflow { what: "key_batch" })?;
if end > payload.len() {
return Err(WireCodecError::Truncated {
what: "key_batch",
needed: len,
at: cursor,
got: payload.len(),
});
}
let key = std::str::from_utf8(&payload[cursor..end])
.map_err(|err| WireCodecError::Utf8 { what: "key_batch", source: err })?
.to_string();
keys.push(key);
cursor = end;
}
if cursor != payload.len() {
return Err(WireCodecError::TrailingBytes {
what: "key_batch",
extra: payload.len() - cursor,
});
}
Ok(keys)
}
#[must_use]
pub fn encode_bytes_batch(items: &[Bytes]) -> Vec<u8> {
let total = BYTES_BATCH_MAGIC.len()
+ 4
+ (items.len() * 4)
+ items.iter().map(Bytes::len).sum::<usize>();
let mut out = Vec::with_capacity(total);
out.extend_from_slice(BYTES_BATCH_MAGIC);
out.extend_from_slice(&(items.len() as u32).to_le_bytes());
for item in items {
out.extend_from_slice(&(item.len() as u32).to_le_bytes());
}
for item in items {
out.extend_from_slice(item);
}
out
}
pub fn decode_bytes_batch(batch: Bytes) -> Result<Vec<Bytes>, WireCodecError> {
if !batch.starts_with(BYTES_BATCH_MAGIC) {
let got_len = BYTES_BATCH_MAGIC.len().min(batch.len());
return Err(WireCodecError::BadMagic {
op: "bytes_batch",
expected: BYTES_BATCH_MAGIC,
got_len,
got: batch[..got_len].to_vec(),
});
}
let data = batch.as_ref();
let mut cursor = BYTES_BATCH_MAGIC.len();
let count = read_u32(data, &mut cursor)? as usize;
let mut lengths = Vec::with_capacity(count);
for _ in 0..count {
lengths.push(read_u32(data, &mut cursor)? as usize);
}
let body_start = cursor;
let body_len = lengths
.iter()
.try_fold(0usize, |acc, len| acc.checked_add(*len))
.ok_or(WireCodecError::LengthOverflow { what: "bytes_batch" })?;
let body_end = body_start
.checked_add(body_len)
.ok_or(WireCodecError::LengthOverflow { what: "bytes_batch" })?;
if body_end != data.len() {
return Err(WireCodecError::BodyLengthMismatch {
what: "bytes_batch",
claimed: body_end,
actual: data.len(),
});
}
let mut offset = body_start;
let mut out = Vec::with_capacity(count);
for len in lengths {
let end = offset + len;
out.push(batch.slice(offset..end));
offset = end;
}
Ok(out)
}
fn read_u32(payload: &[u8], cursor: &mut usize) -> Result<u32, WireCodecError> {
let end = cursor.checked_add(4).ok_or(WireCodecError::LengthOverflow { what: "u32_cursor" })?;
if end > payload.len() {
return Err(WireCodecError::Truncated {
what: "u32",
needed: 4,
at: *cursor,
got: payload.len(),
});
}
let value =
u32::from_le_bytes(payload[*cursor..end].try_into().expect("slice length checked above"));
*cursor = end;
Ok(value)
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone)]
pub struct WireRequest {
pub id: u64,
pub op: u8,
pub namespace: String,
pub key: String,
pub payload: Vec<u8>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone)]
pub struct WireResponse {
pub id: u64,
pub status: u8,
pub op: u8,
pub payload: Vec<u8>,
pub message: String,
}
impl Codec for WireRequest {
type Encoded = AlignedVec;
fn encode(&self) -> Result<Self::Encoded, CodecError> {
rkyv::to_bytes::<RkyvError>(self).map_err(CodecError::encode)
}
fn decode(bytes: &[u8]) -> Result<Self, CodecError> {
let archived = rkyv::access::<<WireRequest as Archive>::Archived, RkyvError>(bytes)
.map_err(CodecError::decode)?;
rkyv::deserialize::<Self, RkyvError>(archived).map_err(CodecError::decode)
}
}
impl Codec for WireResponse {
type Encoded = AlignedVec;
fn encode(&self) -> Result<Self::Encoded, CodecError> {
rkyv::to_bytes::<RkyvError>(self).map_err(CodecError::encode)
}
fn decode(bytes: &[u8]) -> Result<Self, CodecError> {
let archived = rkyv::access::<<WireResponse as Archive>::Archived, RkyvError>(bytes)
.map_err(CodecError::decode)?;
rkyv::deserialize::<Self, RkyvError>(archived).map_err(CodecError::decode)
}
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct LookupBlockPrefixReq {
pub namespace: String,
pub block_hashes_hex: Vec<String>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct LookupBlockPrefixResp {
pub matched_count: u32,
pub error: Option<String>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct GetKvBlocksBatchReq {
pub namespace: String,
pub block_hashes_hex: Vec<String>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct GetKvBlocksBatchResp {
pub payloads: Option<Vec<Vec<u8>>>,
pub error: Option<String>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct PutKvBlocksBatchReq {
pub namespace: String,
pub block_hashes_hex: Vec<String>,
pub payloads: Vec<Vec<u8>>,
}
#[derive(Archive, RkyvSerialize, RkyvDeserialize, Debug, Clone, PartialEq, Eq)]
pub struct PutKvBlocksBatchResp {
pub total_bytes: u64,
pub error: Option<String>,
}
fn expect_magic<'a>(
payload: &'a [u8],
magic: &'static [u8],
op_name: &'static str,
) -> Result<&'a [u8], WireCodecError> {
if !payload.starts_with(magic) {
let got_len = magic.len().min(payload.len());
return Err(WireCodecError::BadMagic {
op: op_name,
expected: magic,
got_len,
got: payload[..got_len].to_vec(),
});
}
Ok(&payload[magic.len()..])
}
fn rkyv_encode_with_magic<T>(
magic: &'static [u8],
op_name: &'static str,
value: &T,
) -> Result<Vec<u8>, WireCodecError>
where
T: for<'a> rkyv::Serialize<
rkyv::api::high::HighSerializer<
rkyv::util::AlignedVec,
rkyv::ser::allocator::ArenaHandle<'a>,
RkyvError,
>,
>,
{
let body = rkyv::to_bytes::<RkyvError>(value)
.map_err(|e| WireCodecError::Rkyv { op: op_name, source: e })?;
let mut out = Vec::with_capacity(magic.len() + body.len());
out.extend_from_slice(magic);
out.extend_from_slice(body.as_slice());
Ok(out)
}
fn rkyv_decode_with_magic<T>(
payload: &[u8],
magic: &'static [u8],
op_name: &'static str,
) -> Result<T, WireCodecError>
where
T: rkyv::Archive,
T::Archived: for<'a> rkyv::bytecheck::CheckBytes<rkyv::api::high::HighValidator<'a, RkyvError>>
+ rkyv::Deserialize<T, rkyv::api::high::HighDeserializer<RkyvError>>,
{
let body = expect_magic(payload, magic, op_name)?;
let mut aligned: rkyv::util::AlignedVec<16> = rkyv::util::AlignedVec::with_capacity(body.len());
aligned.extend_from_slice(body);
rkyv::from_bytes::<T, RkyvError>(&aligned[..])
.map_err(|e| WireCodecError::Rkyv { op: op_name, source: e })
}
pub fn encode_lookup_block_prefix_req(
req: &LookupBlockPrefixReq,
) -> Result<Vec<u8>, WireCodecError> {
rkyv_encode_with_magic(LOOKUP_BLOCK_PREFIX_REQ_MAGIC, "lookup_block_prefix_req", req)
}
pub fn decode_lookup_block_prefix_req(
bytes: &[u8],
) -> Result<LookupBlockPrefixReq, WireCodecError> {
rkyv_decode_with_magic::<LookupBlockPrefixReq>(
bytes,
LOOKUP_BLOCK_PREFIX_REQ_MAGIC,
"lookup_block_prefix_req",
)
}
pub fn encode_lookup_block_prefix_resp(
resp: &LookupBlockPrefixResp,
) -> Result<Vec<u8>, WireCodecError> {
rkyv_encode_with_magic(LOOKUP_BLOCK_PREFIX_RESP_MAGIC, "lookup_block_prefix_resp", resp)
}
pub fn decode_lookup_block_prefix_resp(
bytes: &[u8],
) -> Result<LookupBlockPrefixResp, WireCodecError> {
rkyv_decode_with_magic::<LookupBlockPrefixResp>(
bytes,
LOOKUP_BLOCK_PREFIX_RESP_MAGIC,
"lookup_block_prefix_resp",
)
}
pub fn encode_get_kv_blocks_batch_req(
req: &GetKvBlocksBatchReq,
) -> Result<Vec<u8>, WireCodecError> {
rkyv_encode_with_magic(GET_KV_BLOCKS_BATCH_REQ_MAGIC, "get_kv_blocks_batch_req", req)
}
pub fn decode_get_kv_blocks_batch_req(bytes: &[u8]) -> Result<GetKvBlocksBatchReq, WireCodecError> {
rkyv_decode_with_magic::<GetKvBlocksBatchReq>(
bytes,
GET_KV_BLOCKS_BATCH_REQ_MAGIC,
"get_kv_blocks_batch_req",
)
}
pub fn encode_get_kv_blocks_batch_resp(
resp: &GetKvBlocksBatchResp,
) -> Result<Vec<u8>, WireCodecError> {
rkyv_encode_with_magic(GET_KV_BLOCKS_BATCH_RESP_MAGIC, "get_kv_blocks_batch_resp", resp)
}
pub fn decode_get_kv_blocks_batch_resp(
bytes: &[u8],
) -> Result<GetKvBlocksBatchResp, WireCodecError> {
rkyv_decode_with_magic::<GetKvBlocksBatchResp>(
bytes,
GET_KV_BLOCKS_BATCH_RESP_MAGIC,
"get_kv_blocks_batch_resp",
)
}
pub fn encode_put_kv_blocks_batch_req(
req: &PutKvBlocksBatchReq,
) -> Result<Vec<u8>, WireCodecError> {
if req.block_hashes_hex.len() != req.payloads.len() {
return Err(WireCodecError::BodyLengthMismatch {
what: "put_kv_blocks_batch_req",
claimed: req.block_hashes_hex.len(),
actual: req.payloads.len(),
});
}
rkyv_encode_with_magic(PUT_KV_BLOCKS_BATCH_REQ_MAGIC, "put_kv_blocks_batch_req", req)
}
pub fn decode_put_kv_blocks_batch_req(bytes: &[u8]) -> Result<PutKvBlocksBatchReq, WireCodecError> {
rkyv_decode_with_magic::<PutKvBlocksBatchReq>(
bytes,
PUT_KV_BLOCKS_BATCH_REQ_MAGIC,
"put_kv_blocks_batch_req",
)
}
pub fn encode_put_kv_blocks_batch_resp(
resp: &PutKvBlocksBatchResp,
) -> Result<Vec<u8>, WireCodecError> {
rkyv_encode_with_magic(PUT_KV_BLOCKS_BATCH_RESP_MAGIC, "put_kv_blocks_batch_resp", resp)
}
pub fn decode_put_kv_blocks_batch_resp(
bytes: &[u8],
) -> Result<PutKvBlocksBatchResp, WireCodecError> {
rkyv_decode_with_magic::<PutKvBlocksBatchResp>(
bytes,
PUT_KV_BLOCKS_BATCH_RESP_MAGIC,
"put_kv_blocks_batch_resp",
)
}
const SHM_PREFIX: &str = "wk";
const ROLE_REQ: char = 'r';
const ROLE_RESP: char = 's';
const MAX_DISRUPTOR_INTERNAL_SUFFIX_LEN: usize = 13;
pub fn segment_names(prefix: &str) -> (String, String) {
(format!("{SHM_PREFIX}{prefix}{ROLE_REQ}"), format!("{SHM_PREFIX}{prefix}{ROLE_RESP}"))
}
const SHM_SEGMENT_NAME_MAX_LEN_MACOS: usize = 30;
pub fn validate_segment_name_budget(prefix: &str) -> Result<(), WireCodecError> {
let (req, resp) = segment_names(prefix);
for base in [&req, &resp] {
let derived_len = base.len() + MAX_DISRUPTOR_INTERNAL_SUFFIX_LEN;
if derived_len > SHM_SEGMENT_NAME_MAX_LEN_MACOS {
let fixed = SHM_PREFIX.len() + 1 + MAX_DISRUPTOR_INTERNAL_SUFFIX_LEN;
let max_prefix = SHM_SEGMENT_NAME_MAX_LEN_MACOS.saturating_sub(fixed);
let derived_name = format!("{base}_producer_seq");
return Err(WireCodecError::SegmentNameBudget(format!(
"wombatkv: SHM segment '{derived_name}' ({derived_len} chars) \
exceeds the macOS POSIX-SHM budget of {SHM_SEGMENT_NAME_MAX_LEN_MACOS} \
chars. Shorten the daemon prefix '{prefix}' ({} chars) to at \
most {max_prefix} chars. The internal disruptor-mp segments \
add up to {MAX_DISRUPTOR_INTERNAL_SUFFIX_LEN} chars of suffix \
('_producer_seq' is the longest); the wombatkv wrapper adds \
'{SHM_PREFIX}' + 1-char role. For strict cross-platform \
portability (Linux, FreeBSD), see \
`myelon::portable_shm_segment_name`, recommends total \
name ≤ {} chars.",
prefix.len(),
myelon::PORTABLE_SHM_SEGMENT_NAME_MAX_LEN,
)));
}
}
Ok(())
}
pub fn open_daemon(
req_seg: &str,
resp_seg: &str,
depth: usize,
) -> Result<(TypedConsumer<ShmFrame>, TypedProducer<ShmFrame>), Box<dyn std::error::Error>> {
let req_consumer = wait_for_consumer(req_seg, depth, REQ_CONSUMER_ID)?;
let resp_producer = TypedProducer::<ShmFrame>::create_with_consumers(resp_seg, depth, 1)?;
Ok((req_consumer, resp_producer))
}
pub fn open_client(
req_seg: &str,
resp_seg: &str,
depth: usize,
) -> Result<(TypedProducer<ShmFrame>, TypedConsumer<ShmFrame>), Box<dyn std::error::Error>> {
let req_producer = TypedProducer::<ShmFrame>::create_with_consumers(req_seg, depth, 1)?;
let resp_consumer = wait_for_consumer(resp_seg, depth, RESP_CONSUMER_ID)?;
Ok((req_producer, resp_consumer))
}
fn wait_strategy_from_env() -> MyelonWaitStrategy {
match std::env::var("WMBT_KV_DAEMON_SHM_WAIT_STRATEGY").ok().as_deref().map(str::trim) {
Some("busyspin" | "BusySpin" | "spin") => MyelonWaitStrategy::BusySpin,
Some("block" | "Block" | "park") => MyelonWaitStrategy::Block,
Some(other) if !other.is_empty() => {
eprintln!(
"WombatKV: unknown WMBT_KV_DAEMON_SHM_WAIT_STRATEGY={other:?}, \
falling back to default 'block'"
);
MyelonWaitStrategy::Block
}
_ => MyelonWaitStrategy::Block,
}
}
fn wait_for_consumer(
segment: &str,
depth: usize,
consumer_id: &str,
) -> Result<TypedConsumer<ShmFrame>, Box<dyn std::error::Error>> {
use std::time::Instant;
let deadline = Instant::now() + effective_attach_timeout();
let mut tries: u64 = 0;
let wait_strategy = wait_strategy_from_env();
loop {
tries += 1;
match TypedConsumer::<ShmFrame>::attach_with_consumer_id(
segment,
depth,
consumer_id,
wait_strategy,
) {
Ok(consumer) => return Ok(consumer),
Err(error) => {
if Instant::now() < deadline {
myelon::perform_default_discovery_poll_wait();
continue;
}
return Err(format!("attach {segment}: {error} (after {tries} retries)").into());
}
}
}
}
#[must_use]
pub fn fits_one_frame(len: usize) -> bool {
len.saturating_add(256) <= FRAME_DATA_BYTES
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ping_request_round_trips() {
let req = WireRequest {
id: 42,
op: op::PING,
namespace: String::new(),
key: String::new(),
payload: Vec::new(),
};
let encoded = req.encode().expect("encode");
let decoded = WireRequest::decode(encoded.as_ref()).expect("decode");
assert_eq!(decoded.id, 42);
assert_eq!(decoded.op, op::PING);
}
#[test]
fn put_response_round_trips_with_payload() {
let resp = WireResponse {
id: 7,
status: status::OK,
op: op::PUT,
payload: vec![0xAB; 4096],
message: String::new(),
};
let encoded = resp.encode().expect("encode");
let decoded = WireResponse::decode(encoded.as_ref()).expect("decode");
assert_eq!(decoded.id, 7);
assert_eq!(decoded.status, status::OK);
assert_eq!(decoded.payload.len(), 4096);
assert_eq!(decoded.payload[0], 0xAB);
}
#[test]
fn frame_capacity_check() {
assert!(fits_one_frame(0));
assert!(fits_one_frame(FRAME_DATA_BYTES - 256));
assert!(!fits_one_frame(FRAME_DATA_BYTES));
assert!(!fits_one_frame(FRAME_DATA_BYTES * 2));
}
#[test]
fn key_batch_round_trips() {
let keys = vec!["a".to_string(), "nested/key".to_string()];
let encoded = encode_key_batch(&keys);
assert_eq!(decode_key_batch(&encoded).expect("decode keys"), keys);
}
#[test]
fn bytes_batch_round_trips_without_copying_items() {
let items = vec![Bytes::from_static(b"alpha"), Bytes::new(), Bytes::from_static(b"gamma")];
let encoded = Bytes::from(encode_bytes_batch(&items));
let decoded = decode_bytes_batch(encoded).expect("decode bytes");
assert_eq!(decoded, items);
}
#[test]
fn segment_names_are_unique_per_prefix() {
let (req, resp) = segment_names("smoke");
assert_eq!(req, "wksmoker");
assert_eq!(resp, "wksmokes");
assert_ne!(req, resp);
}
#[test]
fn validate_segment_name_budget_accepts_short_prefix() {
validate_segment_name_budget("drts999").expect("dst-sweep prefix");
validate_segment_name_budget("drts42").expect("dst-sweep prefix");
validate_segment_name_budget("dmcoreds").expect("typical");
}
#[test]
fn validate_segment_name_budget_rejects_too_long_prefix() {
let too_long = "a".repeat(15);
let err = validate_segment_name_budget(&too_long)
.expect_err("15-char prefix must exceed macOS budget");
assert!(err.to_string().contains("exceeds the macOS"));
validate_segment_name_budget(&"a".repeat(14)).expect("14 chars at the edge");
}
#[test]
fn block_opcodes_are_distinct_and_contiguous() {
assert_eq!(op::LOOKUP_BLOCK_PREFIX, 11);
assert_eq!(op::GET_KV_BLOCKS_BATCH, 12);
assert_eq!(op::PUT_KV_BLOCKS_BATCH, 13);
assert_eq!(op::name(op::LOOKUP_BLOCK_PREFIX), "lookup_block_prefix");
assert_eq!(op::name(op::GET_KV_BLOCKS_BATCH), "get_kv_blocks_batch");
assert_eq!(op::name(op::PUT_KV_BLOCKS_BATCH), "put_kv_blocks_batch");
}
#[test]
fn lookup_block_prefix_req_resp_round_trip() {
let req = LookupBlockPrefixReq {
namespace: "ns/alpha".to_string(),
block_hashes_hex: vec!["aa".repeat(32), "bb".repeat(32)],
};
let bytes = encode_lookup_block_prefix_req(&req).expect("encode req");
let decoded = decode_lookup_block_prefix_req(&bytes).expect("decode req");
assert_eq!(decoded.namespace, req.namespace);
assert_eq!(decoded.block_hashes_hex, req.block_hashes_hex);
let resp_ok = LookupBlockPrefixResp { matched_count: 7, error: None };
let bytes_ok = encode_lookup_block_prefix_resp(&resp_ok).expect("encode resp ok");
let decoded_ok = decode_lookup_block_prefix_resp(&bytes_ok).expect("decode resp ok");
assert_eq!(decoded_ok.matched_count, 7);
assert!(decoded_ok.error.is_none());
let resp_err =
LookupBlockPrefixResp { matched_count: 0, error: Some("bad hex at pos 3".to_string()) };
let bytes_err = encode_lookup_block_prefix_resp(&resp_err).expect("encode resp err");
let decoded_err = decode_lookup_block_prefix_resp(&bytes_err).expect("decode resp err");
assert_eq!(decoded_err.matched_count, 0);
assert_eq!(decoded_err.error.as_deref(), Some("bad hex at pos 3"));
}
#[test]
fn get_kv_blocks_batch_req_resp_round_trip() {
let req = GetKvBlocksBatchReq {
namespace: "ns/beta".to_string(),
block_hashes_hex: vec!["11".repeat(32), "22".repeat(32), "33".repeat(32)],
};
let bytes = encode_get_kv_blocks_batch_req(&req).expect("encode req");
let decoded = decode_get_kv_blocks_batch_req(&bytes).expect("decode req");
assert_eq!(decoded.namespace, req.namespace);
assert_eq!(decoded.block_hashes_hex, req.block_hashes_hex);
let payloads = vec![vec![0xAAu8; 4096], vec![0xBBu8; 1024], vec![0xCCu8; 2048]];
let resp_hit = GetKvBlocksBatchResp { payloads: Some(payloads.clone()), error: None };
let bytes_hit = encode_get_kv_blocks_batch_resp(&resp_hit).expect("encode resp hit");
let decoded_hit = decode_get_kv_blocks_batch_resp(&bytes_hit).expect("decode resp hit");
assert!(decoded_hit.error.is_none());
let got = decoded_hit.payloads.expect("payloads present");
assert_eq!(got.len(), 3);
assert_eq!(got[0].len(), 4096);
assert_eq!(got[1].len(), 1024);
assert_eq!(got[2].len(), 2048);
assert_eq!(got[0][0], 0xAA);
let resp_miss = GetKvBlocksBatchResp { payloads: None, error: None };
let bytes_miss = encode_get_kv_blocks_batch_resp(&resp_miss).expect("encode resp miss");
let decoded_miss = decode_get_kv_blocks_batch_resp(&bytes_miss).expect("decode resp miss");
assert!(decoded_miss.payloads.is_none());
}
#[test]
fn put_kv_blocks_batch_req_resp_round_trip() {
let req = PutKvBlocksBatchReq {
namespace: "ns/gamma".to_string(),
block_hashes_hex: vec!["77".repeat(32), "88".repeat(32)],
payloads: vec![vec![0x77u8; 512], vec![0x88u8; 1024]],
};
let bytes = encode_put_kv_blocks_batch_req(&req).expect("encode req");
let decoded = decode_put_kv_blocks_batch_req(&bytes).expect("decode req");
assert_eq!(decoded.namespace, req.namespace);
assert_eq!(decoded.block_hashes_hex, req.block_hashes_hex);
assert_eq!(decoded.payloads.len(), 2);
assert_eq!(decoded.payloads[0], req.payloads[0]);
assert_eq!(decoded.payloads[1].len(), 1024);
let resp_ok = PutKvBlocksBatchResp { total_bytes: 1536, error: None };
let bytes_ok = encode_put_kv_blocks_batch_resp(&resp_ok).expect("encode resp ok");
let decoded_ok = decode_put_kv_blocks_batch_resp(&bytes_ok).expect("decode resp ok");
assert_eq!(decoded_ok.total_bytes, 1536);
assert!(decoded_ok.error.is_none());
let resp_err =
PutKvBlocksBatchResp { total_bytes: 0, error: Some("backend put failure".to_string()) };
let bytes_err = encode_put_kv_blocks_batch_resp(&resp_err).expect("encode resp err");
let decoded_err = decode_put_kv_blocks_batch_resp(&bytes_err).expect("decode resp err");
assert_eq!(decoded_err.total_bytes, 0);
assert_eq!(decoded_err.error.as_deref(), Some("backend put failure"));
}
#[test]
fn block_payload_magic_headers_are_stable() {
assert_eq!(LOOKUP_BLOCK_PREFIX_REQ_MAGIC, b"WMBT_LBP_REQ\0\0\0\0");
assert_eq!(LOOKUP_BLOCK_PREFIX_RESP_MAGIC, b"WMBT_LBP_RES\0\0\0\0");
assert_eq!(GET_KV_BLOCKS_BATCH_REQ_MAGIC, b"WMBT_GBB_REQ\0\0\0\0");
assert_eq!(GET_KV_BLOCKS_BATCH_RESP_MAGIC, b"WMBT_GBB_RES\0\0\0\0");
assert_eq!(PUT_KV_BLOCKS_BATCH_REQ_MAGIC, b"WMBT_PBB_REQ\0\0\0\0");
assert_eq!(PUT_KV_BLOCKS_BATCH_RESP_MAGIC, b"WMBT_PBB_RES\0\0\0\0");
let lbp_req = encode_lookup_block_prefix_req(&LookupBlockPrefixReq {
namespace: "ns".to_string(),
block_hashes_hex: vec!["aa".repeat(32)],
})
.expect("encode");
assert!(lbp_req.starts_with(LOOKUP_BLOCK_PREFIX_REQ_MAGIC));
let gbb_resp = encode_get_kv_blocks_batch_resp(&GetKvBlocksBatchResp {
payloads: Some(vec![vec![0u8; 4]]),
error: None,
})
.expect("encode");
assert!(gbb_resp.starts_with(GET_KV_BLOCKS_BATCH_RESP_MAGIC));
let pbb_resp =
encode_put_kv_blocks_batch_resp(&PutKvBlocksBatchResp { total_bytes: 42, error: None })
.expect("encode");
assert!(pbb_resp.starts_with(PUT_KV_BLOCKS_BATCH_RESP_MAGIC));
}
#[test]
fn block_payload_codec_rejects_corruption() {
let mut bad_magic = encode_lookup_block_prefix_req(&LookupBlockPrefixReq {
namespace: "ns".to_string(),
block_hashes_hex: vec!["aa".repeat(32)],
})
.expect("encode");
bad_magic[0] = b'X';
let err = decode_lookup_block_prefix_req(&bad_magic).expect_err("must fail");
assert!(err.to_string().contains("lookup_block_prefix_req"));
assert!(err.to_string().contains("bad magic"));
let full = encode_put_kv_blocks_batch_req(&PutKvBlocksBatchReq {
namespace: "ns".to_string(),
block_hashes_hex: vec!["aa".repeat(32), "bb".repeat(32)],
payloads: vec![vec![1u8; 8], vec![2u8; 8]],
})
.expect("encode");
let truncated = &full[..full.len() - 8];
let err = decode_put_kv_blocks_batch_req(truncated).expect_err("must fail");
assert!(
err.to_string().contains("put_kv_blocks_batch_req"),
"unexpected error message: {err}"
);
let mut tampered = encode_get_kv_blocks_batch_resp(&GetKvBlocksBatchResp {
payloads: Some(vec![vec![0xAAu8; 8]]),
error: None,
})
.expect("encode");
let mid = usize::midpoint(GET_KV_BLOCKS_BATCH_RESP_MAGIC.len(), tampered.len());
tampered[mid] ^= 0xFF;
let result = decode_get_kv_blocks_batch_resp(&tampered);
if let Ok(decoded) = result {
let _ = decoded;
}
}
#[test]
fn block_payload_encodes_large_blocks_without_overflow() {
let big_block = vec![0xCDu8; 1024 * 1024];
let req = PutKvBlocksBatchReq {
namespace: "ns".to_string(),
block_hashes_hex: (0..4).map(|i| format!("{i:02x}").repeat(32)).collect(),
payloads: vec![big_block.clone(); 4],
};
let bytes = encode_put_kv_blocks_batch_req(&req).expect("encode");
let decoded = decode_put_kv_blocks_batch_req(&bytes).expect("decode");
assert_eq!(decoded.payloads.len(), 4);
for p in &decoded.payloads {
assert_eq!(p.len(), big_block.len());
assert_eq!(p[0], 0xCD);
}
}
}