use std::fmt;
use std::str::FromStr;
use minicbor::decode::{Decoder, Error as DecodeError};
use minicbor::encode::{Encoder, Error as EncodeError, Write};
use minicbor::{Decode, Encode};
use crate::{Cursor, Epoch, Hash, StreamId, StreamName};
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Batch {
#[n(0)]
pub epoch: Epoch,
#[n(1)]
#[cbor(with = "blobs")]
pub entries: Vec<Vec<u8>>,
#[n(2)]
pub stream: Option<StreamId>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Written {
#[n(0)]
pub stream: StreamId,
#[n(1)]
pub epoch: Epoch,
#[n(2)]
pub first: u64,
#[n(3)]
pub last: u64,
}
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Page {
#[n(0)]
pub stream: StreamId,
#[n(1)]
pub epoch: Epoch,
#[n(2)]
pub head: u64,
#[n(3)]
pub entries: Vec<Entry>,
}
impl Page {
pub fn cursor(&self) -> Option<Cursor> {
self.entries.last().map(|entry| entry.cursor(self.stream))
}
}
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Entry {
#[n(0)]
pub seq: u64,
#[n(1)]
pub epoch: Epoch,
#[n(2)]
pub hash: Hash,
#[n(3)]
#[cbor(with = "minicbor::bytes")]
pub data: Vec<u8>,
}
impl Entry {
pub fn cursor(&self, stream: StreamId) -> Cursor {
Cursor { stream, seq: self.seq, hash: self.hash }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Watch {
pub stream: StreamName,
pub cursor: Option<Cursor>,
}
impl fmt::Display for Watch {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.cursor {
Some(cursor) => write!(f, "{}:{cursor}", self.stream),
None => write!(f, "{}", self.stream),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ParseWatchError(String);
impl fmt::Display for ParseWatchError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for ParseWatchError {}
impl FromStr for Watch {
type Err = ParseWatchError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let (name, cursor) = match s.split_once(':') {
Some((name, cursor)) => (name, Some(cursor)),
None => (s, None),
};
let stream = name.parse().map_err(|e: crate::InvalidName| ParseWatchError(e.to_string()))?;
let cursor =
cursor.map(|text| text.parse::<Cursor>().map_err(|e| ParseWatchError(e.to_string()))).transpose()?;
Ok(Self { stream, cursor })
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Heads {
#[n(0)]
pub heads: Vec<Head>,
}
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Head {
#[n(0)]
pub stream: StreamName,
#[n(1)]
pub id: Option<StreamId>,
#[n(2)]
pub head: u64,
}
#[derive(Clone, Debug, PartialEq, Eq, Encode, Decode)]
#[cbor(map)]
pub struct Problem {
#[n(0)]
pub code: ProblemCode,
#[n(1)]
pub message: String,
#[n(2)]
pub head: Option<u64>,
#[n(3)]
pub epoch: Option<Epoch>,
#[n(4)]
pub limit: Option<u64>,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum ProblemCode {
BadRequest,
StreamNotFound,
StreamReplaced,
CursorAhead,
CursorDiverged,
EpochBehind,
TooLarge,
UnsupportedMediaType,
NotFound,
MethodNotAllowed,
Internal,
Other(String),
}
impl ProblemCode {
pub const KNOWN: [ProblemCode; 11] = [
Self::BadRequest,
Self::StreamNotFound,
Self::StreamReplaced,
Self::CursorAhead,
Self::CursorDiverged,
Self::EpochBehind,
Self::TooLarge,
Self::UnsupportedMediaType,
Self::NotFound,
Self::MethodNotAllowed,
Self::Internal,
];
pub fn as_str(&self) -> &str {
match self {
Self::BadRequest => "bad_request",
Self::StreamNotFound => "stream_not_found",
Self::StreamReplaced => "stream_replaced",
Self::CursorAhead => "cursor_ahead",
Self::CursorDiverged => "cursor_diverged",
Self::EpochBehind => "epoch_behind",
Self::TooLarge => "too_large",
Self::UnsupportedMediaType => "unsupported_media_type",
Self::NotFound => "not_found",
Self::MethodNotAllowed => "method_not_allowed",
Self::Internal => "internal",
Self::Other(code) => code,
}
}
pub fn from_wire(code: &str) -> Self {
match code {
"bad_request" => Self::BadRequest,
"stream_not_found" => Self::StreamNotFound,
"stream_replaced" => Self::StreamReplaced,
"cursor_ahead" => Self::CursorAhead,
"cursor_diverged" => Self::CursorDiverged,
"epoch_behind" => Self::EpochBehind,
"too_large" => Self::TooLarge,
"unsupported_media_type" => Self::UnsupportedMediaType,
"not_found" => Self::NotFound,
"method_not_allowed" => Self::MethodNotAllowed,
"internal" => Self::Internal,
other => Self::Other(other.to_owned()),
}
}
}
impl fmt::Display for ProblemCode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl<C> Encode<C> for ProblemCode {
fn encode<W: Write>(&self, e: &mut Encoder<W>, _: &mut C) -> Result<(), EncodeError<W::Error>> {
e.str(self.as_str())?.ok()
}
}
impl<'b, C> Decode<'b, C> for ProblemCode {
fn decode(d: &mut Decoder<'b>, _: &mut C) -> Result<Self, DecodeError> {
Ok(Self::from_wire(d.str()?))
}
}
#[derive(Debug)]
pub struct WireError(DecodeError);
impl fmt::Display for WireError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "not a valid efema message: {}", self.0)
}
}
impl std::error::Error for WireError {}
pub fn encode<T: Encode<()>>(message: &T) -> Vec<u8> {
minicbor::to_vec(message).expect("encoding into a vector is infallible")
}
pub fn decode<'b, T: Decode<'b, ()>>(bytes: &'b [u8]) -> Result<T, WireError> {
let mut decoder = Decoder::new(bytes);
let message = decoder.decode().map_err(WireError)?;
if decoder.position() != bytes.len() {
return Err(WireError(DecodeError::message("trailing bytes after the message").at(decoder.position())));
}
Ok(message)
}
mod blobs {
use minicbor::decode::{Decoder, Error as DecodeError};
use minicbor::encode::{Encoder, Error as EncodeError, Write};
pub fn encode<C, W: Write>(blobs: &[Vec<u8>], e: &mut Encoder<W>, _: &mut C) -> Result<(), EncodeError<W::Error>> {
e.array(blobs.len() as u64)?;
for blob in blobs {
e.bytes(blob)?;
}
Ok(())
}
pub fn decode<C>(d: &mut Decoder<'_>, _: &mut C) -> Result<Vec<Vec<u8>>, DecodeError> {
let position = d.position();
let Some(len) = d.array()? else {
return Err(DecodeError::message("expected a definite-length array of entries").at(position));
};
let mut blobs = Vec::with_capacity(usize::try_from(len).unwrap_or(0).min(d.input().len()));
for _ in 0..len {
blobs.push(d.bytes()?.to_vec());
}
Ok(blobs)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn stream() -> StreamId {
StreamId::from_bytes(*b"0123456789abcdef")
}
#[test]
fn a_batch_round_trips() {
let batch = Batch {
epoch: Epoch(3),
entries: vec![b"one".to_vec(), Vec::new(), vec![0xff; 300]],
stream: Some(stream()),
};
assert_eq!(decode::<Batch>(&encode(&batch)).unwrap(), batch);
}
#[test]
fn entry_data_travels_as_byte_strings() {
let batch = Batch { epoch: Epoch(1), entries: vec![b"hi".to_vec(), Vec::new()], stream: None };
assert_eq!(encode(&batch), [0xa2, 0x00, 0x01, 0x01, 0x82, 0x42, b'h', b'i', 0x40]);
}
#[test]
fn a_batch_names_the_stream_it_expects_as_an_identity() {
let batch = Batch { epoch: Epoch(1), entries: vec![b"x".to_vec()], stream: Some(stream()) };
let mut expected = vec![0xa3, 0x00, 0x01, 0x01, 0x81, 0x41, b'x', 0x02, 0x50];
expected.extend_from_slice(b"0123456789abcdef");
assert_eq!(encode(&batch), expected);
assert_eq!(decode::<Batch>(&expected).unwrap(), batch);
}
#[test]
fn trailing_bytes_are_refused() {
let mut bytes = encode(&Batch { epoch: Epoch(1), entries: vec![b"x".to_vec()], stream: None });
bytes.push(0x00);
assert!(decode::<Batch>(&bytes).is_err());
}
#[test]
fn unknown_keys_are_ignored_and_missing_required_ones_are_not() {
let mut newer = vec![0xa3, 0x00, 0x01, 0x01, 0x81, 0x41, b'x', 0x09, 0x73];
newer.extend_from_slice(b"from a newer client");
assert_eq!(decode::<Batch>(&newer).unwrap().entries, vec![b"x".to_vec()]);
assert!(decode::<Batch>(&[0xa1, 0x00, 0x01]).is_err());
}
#[test]
fn a_claimed_length_is_not_an_allocation() {
let bytes = [0xa2, 0x00, 0x01, 0x01, 0x9a, 0xff, 0xff, 0xff, 0xff];
assert!(decode::<Batch>(&bytes).is_err());
}
#[test]
fn an_indefinite_array_of_entries_is_refused() {
let bytes = [0xa2, 0x00, 0x01, 0x01, 0x9f, 0x41, b'x', 0xff];
assert!(decode::<Batch>(&bytes).is_err());
}
#[test]
fn a_page_names_the_cursor_after_its_last_entry() {
let entry = |seq: u64| Entry { seq, epoch: Epoch(1), hash: Hash::from_bytes([seq as u8; 32]), data: vec![] };
let page = Page { stream: stream(), epoch: Epoch(1), head: 9, entries: vec![entry(4), entry(5)] };
assert_eq!(page.cursor(), Some(Cursor { stream: stream(), seq: 5, hash: Hash::from_bytes([5; 32]) }));
assert_eq!(Page { entries: vec![], ..page.clone() }.cursor(), None);
assert_eq!(decode::<Page>(&encode(&page)).unwrap(), page);
}
#[test]
fn a_problem_keeps_codes_it_does_not_know() {
let problem = Problem {
code: ProblemCode::Other("rate_limited".into()),
message: "slow down".into(),
head: None,
epoch: None,
limit: None,
};
let decoded = decode::<Problem>(&encode(&problem)).unwrap();
assert_eq!(decoded.code.as_str(), "rate_limited");
assert_eq!(ProblemCode::from_wire("cursor_ahead"), ProblemCode::CursorAhead);
}
#[test]
fn every_known_code_round_trips_through_its_spelling() {
for code in ProblemCode::KNOWN {
assert_eq!(ProblemCode::from_wire(code.as_str()), code);
}
}
#[test]
fn the_list_of_known_codes_is_complete() {
for code in ProblemCode::KNOWN {
match code {
ProblemCode::BadRequest
| ProblemCode::StreamNotFound
| ProblemCode::StreamReplaced
| ProblemCode::CursorAhead
| ProblemCode::CursorDiverged
| ProblemCode::EpochBehind
| ProblemCode::TooLarge
| ProblemCode::UnsupportedMediaType
| ProblemCode::NotFound
| ProblemCode::MethodNotAllowed
| ProblemCode::Internal => {}
ProblemCode::Other(_) => panic!("Other is not a code of its own"),
}
}
let distinct: std::collections::HashSet<&str> = ProblemCode::KNOWN.iter().map(ProblemCode::as_str).collect();
assert_eq!(distinct.len(), ProblemCode::KNOWN.len(), "two codes share a spelling");
}
#[test]
fn a_watch_reads_with_and_without_a_cursor() {
let cursor = Cursor::start(stream());
let with: Watch = format!("notes:{cursor}").parse().unwrap();
assert_eq!(with, Watch { stream: "notes".parse().unwrap(), cursor: Some(cursor) });
assert_eq!(with.to_string().parse::<Watch>().unwrap(), with);
assert!("Notes".parse::<Watch>().is_err());
assert!("notes:".parse::<Watch>().is_err());
assert!("notes:12".parse::<Watch>().is_err());
}
}