use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use parking_lot::Mutex;
use crate::distributed::{PeerId, PeerPermissions};
use crate::ipc::{DecodeError, EncodeError, IpcCodec, IpcMessage, IpcSink, IpcSource};
pub trait DataChannel {
type Error;
fn send_frame(&self, frame: Vec<u8>) -> Result<(), Self::Error>;
fn try_recv_frame(&self) -> Result<Option<Vec<u8>>, Self::Error>;
fn is_open(&self) -> bool;
}
#[derive(Debug)]
pub enum WebRtcTransportError<E> {
Channel(E),
Encode(EncodeError),
Decode(DecodeError),
Closed,
}
impl<E: std::fmt::Display> std::fmt::Display for WebRtcTransportError<E> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Channel(e) => write!(f, "data channel error: {e}"),
Self::Encode(e) => write!(f, "frame encode error: {e}"),
Self::Decode(e) => write!(f, "frame decode error: {e}"),
Self::Closed => write!(f, "data channel closed"),
}
}
}
impl<E: std::fmt::Debug + std::fmt::Display> std::error::Error for WebRtcTransportError<E> {}
pub struct WebRtcSink<C> {
channel: C,
permissions: PeerPermissions,
peer: PeerId,
codec: IpcCodec,
}
impl<C> WebRtcSink<C> {
pub fn new(channel: C, permissions: PeerPermissions, peer: PeerId) -> Self {
Self::with_codec(channel, permissions, peer, IpcCodec::Json)
}
pub fn with_codec(
channel: C,
permissions: PeerPermissions,
peer: PeerId,
codec: IpcCodec,
) -> Self {
Self {
channel,
permissions,
peer,
codec,
}
}
pub fn channel(&self) -> &C {
&self.channel
}
pub fn codec(&self) -> IpcCodec {
self.codec
}
}
impl<C: DataChannel> IpcSink for WebRtcSink<C> {
type Error = WebRtcTransportError<C::Error>;
fn send(&mut self, message: &IpcMessage) -> Result<(), Self::Error> {
if !self.channel.is_open() {
return Err(WebRtcTransportError::Closed);
}
let filtered = match message {
IpcMessage::Snapshot(s) => {
IpcMessage::Snapshot(s.filter_readable(&self.permissions, self.peer))
}
IpcMessage::Delta(d) => {
IpcMessage::Delta(d.filter_readable(&self.permissions, self.peer))
}
IpcMessage::CrdtSync(s) => {
IpcMessage::CrdtSync(s.filter_readable(&self.permissions, self.peer))
}
control @ (IpcMessage::ResyncRequest(_)
| IpcMessage::OutboxAck(_)
| IpcMessage::DeltaSinceRequest(_)) => control.clone(),
};
let frame = self
.codec
.encode(&filtered)
.map_err(WebRtcTransportError::Encode)?;
self.channel
.send_frame(frame)
.map_err(WebRtcTransportError::Channel)
}
}
pub struct WebRtcSource<C> {
channel: C,
codec: IpcCodec,
}
impl<C> WebRtcSource<C> {
pub fn new(channel: C) -> Self {
Self::with_codec(channel, IpcCodec::Json)
}
pub fn with_codec(channel: C, codec: IpcCodec) -> Self {
Self { channel, codec }
}
pub fn channel(&self) -> &C {
&self.channel
}
pub fn codec(&self) -> IpcCodec {
self.codec
}
}
impl<C: DataChannel> IpcSource for WebRtcSource<C> {
type Error = WebRtcTransportError<C::Error>;
fn recv(&mut self) -> Result<Option<IpcMessage>, Self::Error> {
match self
.channel
.try_recv_frame()
.map_err(WebRtcTransportError::Channel)?
{
Some(frame) => Ok(Some(
self.codec
.decode(&frame)
.map_err(WebRtcTransportError::Decode)?,
)),
None => {
if self.channel.is_open() {
Ok(None)
} else {
Err(WebRtcTransportError::Closed)
}
}
}
}
}
#[derive(Clone)]
pub struct InMemoryDataChannel {
tx: Arc<Mutex<VecDeque<Vec<u8>>>>,
rx: Arc<Mutex<VecDeque<Vec<u8>>>>,
open: Arc<AtomicBool>,
}
impl InMemoryDataChannel {
pub fn pair() -> (Self, Self) {
let a_to_b = Arc::new(Mutex::new(VecDeque::new()));
let b_to_a = Arc::new(Mutex::new(VecDeque::new()));
let open = Arc::new(AtomicBool::new(true));
let a = Self {
tx: a_to_b.clone(),
rx: b_to_a.clone(),
open: open.clone(),
};
let b = Self {
tx: b_to_a,
rx: a_to_b,
open,
};
(a, b)
}
pub fn close(&self) {
self.open.store(false, Ordering::SeqCst);
}
}
impl DataChannel for InMemoryDataChannel {
type Error = std::convert::Infallible;
fn send_frame(&self, frame: Vec<u8>) -> Result<(), Self::Error> {
self.tx.lock().push_back(frame);
Ok(())
}
fn try_recv_frame(&self) -> Result<Option<Vec<u8>>, Self::Error> {
Ok(self.rx.lock().pop_front())
}
fn is_open(&self) -> bool {
self.open.load(Ordering::SeqCst)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::distributed::{NodeId, PeerId, PeerPermissions, RemoteOp};
use crate::ipc::{IpcMessage, NodeSnapshot, Snapshot};
fn snapshot_two_nodes() -> Snapshot {
Snapshot::new(
1,
vec![
NodeSnapshot::payload(NodeId(1), "t", vec![1, 2, 3]),
NodeSnapshot::payload(NodeId(2), "t", vec![4, 5, 6]),
],
vec![],
vec![NodeId(1), NodeId(2)],
)
}
#[test]
fn loopback_round_trips_and_filters_unreadable_nodes() {
let (here, there) = InMemoryDataChannel::pair();
let peer = PeerId(7);
let mut perms = PeerPermissions::new();
perms.allow(peer, RemoteOp::read(NodeId(1)));
let mut sink = WebRtcSink::new(here, perms, peer);
let mut source = WebRtcSource::new(there);
sink.send(&IpcMessage::Snapshot(snapshot_two_nodes()))
.unwrap();
let received = source.recv().unwrap().expect("a message");
match received {
IpcMessage::Snapshot(s) => {
let ids: Vec<u64> = s.nodes.iter().map(|n| n.node.0).collect();
assert_eq!(ids, vec![1], "node 2 must be filtered out for this peer");
}
other => panic!("expected snapshot, got {other:?}"),
}
assert!(source.recv().unwrap().is_none());
}
#[test]
fn closed_channel_reports_closed() {
let (here, there) = InMemoryDataChannel::pair();
here.close();
let mut sink = WebRtcSink::new(here, PeerPermissions::new(), PeerId(1));
assert!(matches!(
sink.send(&IpcMessage::Snapshot(snapshot_two_nodes())),
Err(WebRtcTransportError::Closed)
));
let mut source = WebRtcSource::new(there);
assert!(matches!(source.recv(), Err(WebRtcTransportError::Closed)));
}
}