use std::sync::Arc;
use anyhow::Result;
use async_trait::async_trait;
use bytes::Bytes;
use super::demux::{RelayPacketKind, classify_relay_packet};
use super::transport::{RelayTransport, RelayTransportEvent, RelayTransportFactory};
use crate::runtime::Runtime;
const TAP_FORWARD_CAP: usize = 256;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum PacketDir {
Outbound,
Inbound,
}
pub(crate) trait PacketTap: crate::sync_marker::MaybeSendSync {
fn on_packet(&self, dir: PacketDir, data: &[u8]);
}
pub(crate) struct TappedTransport {
inner: Arc<dyn RelayTransport>,
tap: Arc<dyn PacketTap>,
}
impl TappedTransport {
pub fn new(inner: Arc<dyn RelayTransport>, tap: Arc<dyn PacketTap>) -> Self {
Self { inner, tap }
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl RelayTransport for TappedTransport {
async fn send(&self, data: Bytes) -> Result<()> {
self.tap.on_packet(PacketDir::Outbound, &data);
self.inner.send(data).await
}
async fn disconnect(&self) {
self.inner.disconnect().await;
}
}
async fn tap_forward(
inner_rx: async_channel::Receiver<RelayTransportEvent>,
out_tx: async_channel::Sender<RelayTransportEvent>,
tap: Arc<dyn PacketTap>,
) {
while let Ok(ev) = inner_rx.recv().await {
if let RelayTransportEvent::PacketReceived(data) = &ev {
tap.on_packet(PacketDir::Inbound, data);
}
match out_tx.try_send(ev) {
Ok(()) => {}
Err(async_channel::TrySendError::Full(ev)) => {
let is_stun = matches!(&ev, RelayTransportEvent::PacketReceived(d)
if classify_relay_packet(d) == RelayPacketKind::Stun);
if is_stun && out_tx.send(ev).await.is_err() {
break;
}
}
Err(async_channel::TrySendError::Closed(_)) => break,
}
}
}
pub(crate) struct TappedFactory {
inner: Arc<dyn RelayTransportFactory>,
tap: Arc<dyn PacketTap>,
runtime: Arc<dyn Runtime>,
}
impl TappedFactory {
pub fn new(
inner: Arc<dyn RelayTransportFactory>,
tap: Arc<dyn PacketTap>,
runtime: Arc<dyn Runtime>,
) -> Self {
Self {
inner,
tap,
runtime,
}
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl RelayTransportFactory for TappedFactory {
async fn connect(
&self,
) -> Result<(
Arc<dyn RelayTransport>,
async_channel::Receiver<RelayTransportEvent>,
)> {
let (inner_transport, inner_rx) = self.inner.connect().await?;
let transport: Arc<dyn RelayTransport> =
Arc::new(TappedTransport::new(inner_transport, self.tap.clone()));
let (out_tx, out_rx) = async_channel::bounded(TAP_FORWARD_CAP);
self.runtime
.spawn(Box::pin(tap_forward(inner_rx, out_tx, self.tap.clone())))
.detach();
Ok((transport, out_rx))
}
}
#[derive(Default)]
pub(crate) struct InMemoryTap {
captured: std::sync::Mutex<Vec<(PacketDir, Vec<u8>)>>,
}
impl InMemoryTap {
pub fn captured(&self) -> Vec<(PacketDir, Vec<u8>)> {
self.captured.lock().unwrap().clone()
}
}
impl PacketTap for InMemoryTap {
fn on_packet(&self, dir: PacketDir, data: &[u8]) {
self.captured.lock().unwrap().push((dir, data.to_vec()));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::voip::RelayDisconnectReason;
use std::sync::Mutex;
#[derive(Default)]
struct RecordingTransport {
sent: Mutex<Vec<Bytes>>,
}
#[async_trait]
impl RelayTransport for RecordingTransport {
async fn send(&self, data: Bytes) -> Result<()> {
self.sent.lock().unwrap().push(data);
Ok(())
}
async fn disconnect(&self) {}
}
#[test]
fn tapped_transport_records_outbound_then_delegates() {
let inner = Arc::new(RecordingTransport::default());
let tap = Arc::new(InMemoryTap::default());
let tapped = TappedTransport::new(inner.clone(), tap.clone());
futures::executor::block_on(async {
tapped.send(Bytes::from_static(b"\x01\x02")).await.unwrap();
tapped.send(Bytes::from_static(b"\x03")).await.unwrap();
});
assert_eq!(
tap.captured(),
vec![
(PacketDir::Outbound, vec![1, 2]),
(PacketDir::Outbound, vec![3]),
]
);
assert_eq!(inner.sent.lock().unwrap().len(), 2);
}
#[test]
fn tap_forward_records_inbound_and_forwards_every_event() {
let (inner_tx, inner_rx) = async_channel::unbounded();
let (out_tx, out_rx) = async_channel::unbounded();
inner_tx
.try_send(RelayTransportEvent::PacketReceived(Bytes::from_static(
b"\xaa",
)))
.unwrap();
inner_tx.try_send(RelayTransportEvent::Connected).unwrap();
inner_tx
.try_send(RelayTransportEvent::PacketReceived(Bytes::from_static(
b"\xbb\xcc",
)))
.unwrap();
inner_tx
.try_send(RelayTransportEvent::Disconnected(
RelayDisconnectReason::Closed,
))
.unwrap();
inner_tx.close();
let tap = Arc::new(InMemoryTap::default());
futures::executor::block_on(tap_forward(inner_rx, out_tx, tap.clone()));
assert_eq!(
tap.captured(),
vec![
(PacketDir::Inbound, vec![0xaa]),
(PacketDir::Inbound, vec![0xbb, 0xcc]),
]
);
let forwarded: Vec<_> = std::iter::from_fn(|| out_rx.try_recv().ok()).collect();
assert_eq!(forwarded.len(), 4);
assert!(matches!(forwarded[1], RelayTransportEvent::Connected));
assert!(matches!(forwarded[3], RelayTransportEvent::Disconnected(_)));
}
#[test]
fn tap_forward_preserves_stun_but_drops_media_under_backpressure() {
let (inner_tx, inner_rx) = async_channel::unbounded();
for pkt in [
&b"\x90\x78\x01\x02"[..], &b"\x90\x78\x03\x04"[..], &b"\x00\x01\x05\x06"[..], ] {
inner_tx
.try_send(RelayTransportEvent::PacketReceived(Bytes::copy_from_slice(
pkt,
)))
.unwrap();
}
inner_tx.close();
let (out_tx, out_rx) = async_channel::bounded(1);
let tap = Arc::new(InMemoryTap::default());
futures::executor::block_on(async {
let fwd = tap_forward(inner_rx, out_tx, tap.clone());
let drain = async {
let a = out_rx.recv().await.unwrap();
let b = out_rx.recv().await.unwrap();
(a, b)
};
let (_, (a, b)) = futures::join!(fwd, drain);
assert!(
matches!(&a, RelayTransportEvent::PacketReceived(d) if d[0] == 0x90),
"first delivered is the media that filled the channel, got {a:?}"
);
assert!(
matches!(&b, RelayTransportEvent::PacketReceived(d)
if classify_relay_packet(d) == RelayPacketKind::Stun),
"STUN must survive while the media behind the first was dropped, got {b:?}"
);
});
assert_eq!(
tap.captured().len(),
3,
"the tap records every packet, even ones later dropped"
);
}
}