pub mod frame {
pub const SYNC_CONTROL_TAG: &[u8] = b"haematite.sync.v1";
const FRAME_HEADER_BYTES: usize = 8;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum FrameError {
TooShort,
UnexpectedControlTag,
ControlLengthMismatch,
}
impl core::fmt::Display for FrameError {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::TooShort => formatter.write_str("buffer is shorter than a sync frame header"),
Self::UnexpectedControlTag => {
formatter.write_str("binary frame is not haematite sync-protocol traffic")
}
Self::ControlLengthMismatch => {
formatter.write_str("sync frame control-tag length is malformed")
}
}
}
}
impl core::error::Error for FrameError {}
fn read_u32_be(bytes: &[u8], offset: usize) -> Option<u32> {
let end = offset.checked_add(4)?;
let slice = bytes.get(offset..end)?;
let mut value = [0_u8; 4];
value.copy_from_slice(slice);
Some(u32::from_be_bytes(value))
}
pub fn validate_sync_frame(bytes: &[u8]) -> Result<(), FrameError> {
if bytes.len() < FRAME_HEADER_BYTES {
return Err(FrameError::TooShort);
}
let control_len = read_u32_be(bytes, 0).ok_or(FrameError::TooShort)? as usize;
if control_len != SYNC_CONTROL_TAG.len() {
return Err(FrameError::ControlLengthMismatch);
}
let tag_start = FRAME_HEADER_BYTES;
let tag_end = tag_start
.checked_add(control_len)
.ok_or(FrameError::TooShort)?;
let tag = bytes.get(tag_start..tag_end).ok_or(FrameError::TooShort)?;
if tag != SYNC_CONTROL_TAG {
return Err(FrameError::UnexpectedControlTag);
}
Ok(())
}
pub fn outbound_binary_payload(frame: &[u8]) -> Result<&[u8], FrameError> {
validate_sync_frame(frame)?;
Ok(frame)
}
pub fn inbound_frame(payload: Vec<u8>) -> Result<Vec<u8>, FrameError> {
validate_sync_frame(&payload)?;
Ok(payload)
}
}
#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
pub use socket::{WebSocketSyncTransport, WebSocketTransportError};
#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
mod socket {
use std::cell::RefCell;
use std::collections::VecDeque;
use std::rc::Rc;
use js_sys::{ArrayBuffer, Uint8Array};
use wasm_bindgen::JsCast;
use wasm_bindgen::closure::Closure;
use wasm_bindgen::prelude::*;
use web_sys::{BinaryType, MessageEvent, WebSocket};
use super::frame;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WebSocketTransportError {
Open(String),
Send(String),
Frame(frame::FrameError),
NonBinaryMessage,
}
impl core::fmt::Display for WebSocketTransportError {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Open(message) => write!(formatter, "websocket open failed: {message}"),
Self::Send(message) => write!(formatter, "websocket send failed: {message}"),
Self::Frame(error) => write!(formatter, "websocket frame error: {error}"),
Self::NonBinaryMessage => {
formatter.write_str("websocket message was text, expected a binary sync frame")
}
}
}
}
impl core::error::Error for WebSocketTransportError {}
impl From<frame::FrameError> for WebSocketTransportError {
fn from(error: frame::FrameError) -> Self {
Self::Frame(error)
}
}
fn js_message(value: &JsValue) -> String {
value
.as_string()
.or_else(|| js_sys::Error::from(value.clone()).message().as_string())
.unwrap_or_else(|| String::from("unknown JavaScript error"))
}
pub struct WebSocketSyncTransport {
socket: WebSocket,
inbound: Rc<RefCell<VecDeque<Vec<u8>>>>,
_on_message: Closure<dyn FnMut(MessageEvent)>,
}
impl WebSocketSyncTransport {
pub fn connect(endpoint: &str) -> Result<Self, WebSocketTransportError> {
let socket = WebSocket::new(endpoint)
.map_err(|error| WebSocketTransportError::Open(js_message(&error)))?;
socket.set_binary_type(BinaryType::Arraybuffer);
let inbound: Rc<RefCell<VecDeque<Vec<u8>>>> = Rc::new(RefCell::new(VecDeque::new()));
let inbound_for_cb = Rc::clone(&inbound);
let on_message = Closure::wrap(Box::new(move |event: MessageEvent| {
let data = event.data();
if let Some(buffer) = data.dyn_ref::<ArrayBuffer>() {
let bytes = Uint8Array::new(buffer).to_vec();
if let Ok(valid) = frame::inbound_frame(bytes) {
inbound_for_cb.borrow_mut().push_back(valid);
}
}
}) as Box<dyn FnMut(MessageEvent)>);
socket.set_onmessage(Some(on_message.as_ref().unchecked_ref()));
Ok(Self {
socket,
inbound,
_on_message: on_message,
})
}
pub fn send_frame(&self, frame_bytes: &[u8]) -> Result<(), WebSocketTransportError> {
let payload = frame::outbound_binary_payload(frame_bytes)?;
self.socket
.send_with_u8_array(payload)
.map_err(|error| WebSocketTransportError::Send(js_message(&error)))
}
pub fn start_pull_sync(
&self,
pull_request_frame: &[u8],
) -> Result<(), WebSocketTransportError> {
self.send_frame(pull_request_frame)
}
pub fn drain_inbound(&self) -> Vec<Vec<u8>> {
self.inbound.borrow_mut().drain(..).collect()
}
pub fn pending_inbound(&self) -> usize {
self.inbound.borrow().len()
}
pub fn ready_state(&self) -> u16 {
self.socket.ready_state()
}
pub const fn socket(&self) -> &WebSocket {
&self.socket
}
}
}
#[cfg(test)]
mod tests {
use super::frame;
use crate::sync_codec::ballot::Ballot;
use crate::sync_codec::{
NodeTransfer, PullRequest, PushResponse, RootExchangeRequest, SyncMessage, SyncStats,
WriteId, WriteProposal, decode_beamr_sync_frame, encode_beamr_sync_frame,
};
use crate::tree::{LeafNode, Node};
type TestResult = Result<(), Box<dyn std::error::Error>>;
fn leaf(key: &[u8], value: &[u8]) -> Result<Node, Box<dyn std::error::Error>> {
Ok(Node::Leaf(LeafNode::new(vec![(
key.to_vec(),
value.to_vec(),
)])?))
}
fn sample_messages() -> Result<Vec<SyncMessage>, Box<dyn std::error::Error>> {
let transfer = NodeTransfer::new(leaf(b"alpha", b"one")?);
let push = PushResponse::new(0, None, None, vec![transfer], SyncStats::default());
Ok(vec![
SyncMessage::RootRequest(RootExchangeRequest::new(3, None)),
SyncMessage::PullRequest(PullRequest::new(1, None)),
SyncMessage::PushResponse(push),
SyncMessage::WriteProposal(WriteProposal {
write_id: WriteId::new("node-a", 7, 42),
shard_id: 2,
key: b"k".to_vec(),
expected: None,
value: b"v".to_vec(),
ttl: None,
epoch: Ballot::bottom(),
seq: 0,
tombstone: false,
}),
])
}
#[test]
fn outbound_payload_is_the_codec_frame_verbatim() -> TestResult {
for message in sample_messages()? {
let frame_bytes = encode_beamr_sync_frame(&message)?;
let payload = frame::outbound_binary_payload(&frame_bytes)?;
assert_eq!(
payload, frame_bytes,
"carrier must ship the exact codec frame bytes"
);
}
Ok(())
}
#[test]
fn frame_round_trips_through_carrier_to_identical_message() -> TestResult {
for message in sample_messages()? {
let frame_bytes = encode_beamr_sync_frame(&message)?;
let on_wire = frame::outbound_binary_payload(&frame_bytes)?.to_vec();
let received = frame::inbound_frame(on_wire)?;
let decoded = decode_beamr_sync_frame(&received)?;
assert_eq!(decoded, message, "round-trip must preserve the message");
}
Ok(())
}
#[test]
fn wasm_wire_bytes_match_native_codec_bytes() -> TestResult {
let message = SyncMessage::PushResponse(PushResponse::new(
0,
None,
None,
vec![NodeTransfer::new(leaf(b"k", b"v")?)],
SyncStats::default(),
));
let native_frame = encode_beamr_sync_frame(&message)?;
let wasm_wire = frame::outbound_binary_payload(&native_frame)?.to_vec();
assert_eq!(wasm_wire, native_frame);
let decoded = decode_beamr_sync_frame(&wasm_wire)?;
assert_eq!(decoded, message);
Ok(())
}
#[test]
fn validate_rejects_non_sync_frames() {
assert_eq!(
frame::validate_sync_frame(&[]),
Err(frame::FrameError::TooShort)
);
assert_eq!(
frame::validate_sync_frame(&[0, 0, 0, 0, 0, 0, 0, 0]),
Err(frame::FrameError::ControlLengthMismatch)
);
let mut wrong_tag = Vec::new();
wrong_tag.extend_from_slice(&(frame::SYNC_CONTROL_TAG.len() as u32).to_be_bytes());
wrong_tag.extend_from_slice(&0_u32.to_be_bytes());
wrong_tag.extend_from_slice(b"not-a-sync-tag-xx");
assert_eq!(
frame::validate_sync_frame(&wrong_tag),
Err(frame::FrameError::UnexpectedControlTag)
);
}
#[test]
fn validate_accepts_a_real_codec_frame() -> TestResult {
let frame_bytes =
encode_beamr_sync_frame(&SyncMessage::PullRequest(PullRequest::new(0, None)))?;
assert_eq!(frame::validate_sync_frame(&frame_bytes), Ok(()));
Ok(())
}
}