#![cfg(all(not(target_arch = "wasm32"), feature = "iroh-transport-webrtc"))]
use std::{
sync::{
atomic::{AtomicBool, Ordering},
Arc, Mutex,
},
time::Instant,
};
use bytes::Bytes;
use tokio::task::JoinHandle;
use crate::{
iroh_carrier::{
segment_packet, CarrierControl, CarrierFrame, CarrierFrameExpectation, CarrierReassembler,
CARRIER_CONTROL_PACKET_ID,
},
packet_carrier_transport::PacketCarrierSession,
transport::WebRtcDataChannel,
};
const DATA_CHANNEL_MESSAGE_CEILING: usize = 16 * 1024;
const BUFFERED_AMOUNT_HARD_LIMIT: usize = 1024 * 1024;
pub struct NativeWebRtcCarrierSession {
channel: Arc<WebRtcDataChannel>,
packet_session: Arc<PacketCarrierSession>,
outbound_task: JoinHandle<()>,
expected: CarrierFrameExpectation,
application_key: [u8; 32],
terminal_reason: Arc<Mutex<Option<&'static str>>>,
terminal_acknowledged: Arc<AtomicBool>,
}
impl std::fmt::Debug for NativeWebRtcCarrierSession {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("NativeWebRtcCarrierSession")
.finish_non_exhaustive()
}
}
impl NativeWebRtcCarrierSession {
pub async fn attach(
channel: Arc<WebRtcDataChannel>,
session: PacketCarrierSession,
expected: CarrierFrameExpectation,
application_key: [u8; 32],
) -> anyhow::Result<Self> {
let session = Arc::new(session);
let reassembler = Arc::new(tokio::sync::Mutex::new(CarrierReassembler::default()));
let reassembly_clock = Instant::now();
let inbound_session = session.clone();
let inbound_reassembler = reassembler;
let inbound_channel = channel.clone();
let terminal_reason = Arc::new(Mutex::new(None));
let terminal_acknowledged = Arc::new(AtomicBool::new(false));
let inbound_terminal_reason = terminal_reason.clone();
let inbound_terminal_acknowledged = terminal_acknowledged.clone();
channel
.set_message_handler(Arc::new(move |message: Bytes| {
let session = inbound_session.clone();
let reassembler = inbound_reassembler.clone();
let channel = inbound_channel.clone();
let terminal_reason = inbound_terminal_reason.clone();
let terminal_acknowledged = inbound_terminal_acknowledged.clone();
Box::pin(async move {
let Ok(frame) = CarrierFrame::decode(&message, expected) else {
return;
};
match frame.terminal_control_kind(expected, &application_key) {
Ok(Some(CarrierControl::SessionTokenRevokedAck)) => {
terminal_acknowledged.store(true, Ordering::Release);
return;
}
Ok(Some(control @ CarrierControl::SessionTokenRevoked)) => {
if let Ok(ack) = CarrierFrame::terminal_control(
expected,
&application_key,
CarrierControl::SessionTokenRevokedAck,
)
.encode()
{
let _ = channel.send(&ack).await;
}
if let (Some(lifecycle_reason), Ok(mut reason)) =
(control.lifecycle_reason(), terminal_reason.lock())
{
*reason = Some(lifecycle_reason);
}
session.close();
return;
}
Err(_) if frame.header.packet_id == CARRIER_CONTROL_PACKET_ID => return,
_ => {}
}
let now_ms =
u64::try_from(reassembly_clock.elapsed().as_millis()).unwrap_or(u64::MAX);
let Ok(Some(packet)) = reassembler.lock().await.push(frame, now_ms) else {
return;
};
let _ = session.deliver_inbound(packet).await;
})
}))
.await;
let outbound_channel = channel.clone();
let outbound_session = session.clone();
let outbound_task = tokio::spawn(async move {
let mut packet_id = 0_u64;
while let Ok(packet) = outbound_session.recv_outbound().await {
while outbound_channel.buffered_amount().await > BUFFERED_AMOUNT_HARD_LIMIT {
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
}
packet_id = packet_id.wrapping_add(1);
if packet_id == CARRIER_CONTROL_PACKET_ID {
packet_id = 0;
}
let Ok(frames) =
segment_packet(&packet, expected, packet_id, DATA_CHANNEL_MESSAGE_CEILING)
else {
break;
};
for frame in frames {
let Ok(encoded) = frame.encode() else {
return;
};
if outbound_channel.send(&encoded).await.is_err() {
return;
}
}
}
});
Ok(Self {
channel,
packet_session: session,
outbound_task,
expected,
application_key,
terminal_reason,
terminal_acknowledged,
})
}
pub fn take_terminal_reason(&self) -> Option<&'static str> {
self.terminal_reason.lock().ok()?.take()
}
pub async fn send_terminal(&self, reason: &str) -> anyhow::Result<()> {
let control = match reason {
crate::lifecycle_reason::REASON_SESSION_TOKEN_REVOKED => {
CarrierControl::SessionTokenRevoked
}
_ => return Ok(()),
};
let encoded = CarrierFrame::terminal_control(self.expected, &self.application_key, control)
.encode()?;
self.terminal_acknowledged.store(false, Ordering::Release);
for _ in 0..12 {
self.channel.send(&encoded).await?;
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
if self.terminal_acknowledged.load(Ordering::Acquire) {
break;
}
}
Ok(())
}
}
impl Drop for NativeWebRtcCarrierSession {
fn drop(&mut self) {
self.packet_session.close();
self.outbound_task.abort();
self.channel.close();
}
}
#[cfg(all(test, feature = "iroh-protocols-wasm"))]
mod tests {
use super::*;
use crate::{
client::{TransportImpl, WebRTCConfig},
iroh_carrier_kind::EXPERIMENTAL_WEBRTC_TRANSPORT_ID,
packet_carrier_transport::PacketCarrierTransport,
transport::{NativeWebRTCRole, NativeWebRTCState, WebRTCSignalMessage},
};
use iroh::{endpoint::presets, Endpoint, EndpointAddr, RelayMode, TransportAddr};
use tokio::sync::mpsc;
const TEST_ALPN: &[u8] = b"openrtc/test/webrtc-iroh-carrier/1";
const TEST_SESSION_ID: [u8; 16] = [0x7b; 16];
fn parse_signal_frame(frame: serde_json::Value) -> Option<WebRTCSignalMessage> {
if frame.get("type").and_then(serde_json::Value::as_str) != Some("#pluto-signal") {
return None;
}
let content = frame.get("content")?;
Some(WebRTCSignalMessage {
transport: content
.get("transport")
.and_then(serde_json::Value::as_str)
.unwrap_or("webrtc")
.to_string(),
signal_type: content
.get("type")
.and_then(serde_json::Value::as_str)
.unwrap_or("candidate")
.to_string(),
negotiation_id: content
.get("negotiationId")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
sdp: content.get("sdp").cloned(),
candidate: content.get("candidate").cloned(),
})
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn two_iroh_endpoints_cross_negotiated_webrtc_data_channel() {
let secret_a = iroh::SecretKey::generate();
let secret_b = iroh::SecretKey::generate();
let endpoint_id_a = secret_a.public();
let endpoint_id_b = secret_b.public();
let transport_a = Arc::new(
PacketCarrierTransport::new(EXPERIMENTAL_WEBRTC_TRANSPORT_ID, endpoint_id_a).unwrap(),
);
let transport_b = Arc::new(
PacketCarrierTransport::new(EXPERIMENTAL_WEBRTC_TRANSPORT_ID, endpoint_id_b).unwrap(),
);
let packet_session_a = transport_a
.attach_session(transport_a.peer_addr(endpoint_id_b))
.unwrap();
let remote_addr_b = packet_session_a.remote_addr().clone();
let packet_session_b = transport_b
.attach_session(transport_b.peer_addr(endpoint_id_a))
.unwrap();
let config = WebRTCConfig {
implementation: TransportImpl::IrohCarrier,
ice_servers: Vec::new(),
privacy_mode: false,
lan_mode: true,
};
let (a_to_b_tx, mut a_to_b_rx) = mpsc::unbounded_channel();
let (b_to_a_tx, mut b_to_a_rx) = mpsc::unbounded_channel();
let sender_a: crate::transport::WebRtcSignalSender = Arc::new(move |signal| {
let tx = a_to_b_tx.clone();
Box::pin(async move {
let _ = tx.send(signal);
})
});
let sender_b: crate::transport::WebRtcSignalSender = Arc::new(move |signal| {
let tx = b_to_a_tx.clone();
Box::pin(async move {
let _ = tx.send(signal);
})
});
let channel_a = Arc::new(
WebRtcDataChannel::new_packet_carrier(
&endpoint_id_a.to_string(),
&endpoint_id_b.to_string(),
&config,
sender_a,
"carrier-loopback",
NativeWebRTCRole::Initiator,
)
.await
.unwrap(),
);
let channel_b = Arc::new(
WebRtcDataChannel::new_packet_carrier(
&endpoint_id_b.to_string(),
&endpoint_id_a.to_string(),
&config,
sender_b,
"carrier-loopback",
NativeWebRTCRole::Responder,
)
.await
.unwrap(),
);
let expected = CarrierFrameExpectation {
transport_id: EXPERIMENTAL_WEBRTC_TRANSPORT_ID,
carrier_session_id: TEST_SESSION_ID,
transport_generation: 5,
};
let application_key = [0x5c; 32];
let carrier_a = NativeWebRtcCarrierSession::attach(
channel_a.clone(),
packet_session_a,
expected,
application_key,
)
.await
.unwrap();
let carrier_b = NativeWebRtcCarrierSession::attach(
channel_b.clone(),
packet_session_b,
expected,
application_key,
)
.await
.unwrap();
let channel_b_for_signals = channel_b.clone();
let relay_a_to_b = tokio::spawn(async move {
while let Some(frame) = a_to_b_rx.recv().await {
if let Some(signal) = parse_signal_frame(frame) {
let _ = channel_b_for_signals.handle_signal(signal).await;
}
}
});
let channel_a_for_signals = channel_a.clone();
let relay_b_to_a = tokio::spawn(async move {
while let Some(frame) = b_to_a_rx.recv().await {
if let Some(signal) = parse_signal_frame(frame) {
let _ = channel_a_for_signals.handle_signal(signal).await;
}
}
});
channel_b.start().await.unwrap();
channel_a.start().await.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(10), async {
loop {
if channel_a.state() == NativeWebRTCState::Connected
&& channel_b.state() == NativeWebRTCState::Connected
{
break;
}
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
}
})
.await
.expect("negotiated WebRTC carrier should connect");
let endpoint_a = Endpoint::builder(presets::N0)
.secret_key(secret_a)
.relay_mode(RelayMode::Disabled)
.clear_ip_transports()
.alpns(vec![
TEST_ALPN.to_vec(),
iroh_blobs::ALPN.to_vec(),
iroh_gossip::ALPN.to_vec(),
iroh_docs::ALPN.to_vec(),
])
.add_custom_transport(transport_a)
.bind()
.await
.unwrap();
let endpoint_b = Endpoint::builder(presets::N0)
.secret_key(secret_b)
.relay_mode(RelayMode::Disabled)
.clear_ip_transports()
.alpns(vec![
TEST_ALPN.to_vec(),
iroh_blobs::ALPN.to_vec(),
iroh_gossip::ALPN.to_vec(),
iroh_docs::ALPN.to_vec(),
])
.add_custom_transport(transport_b)
.bind()
.await
.unwrap();
let server = {
let endpoint_b = endpoint_b.clone();
tokio::spawn(async move {
let incoming = endpoint_b.accept().await.expect("incoming WebRTC carrier");
let connection = incoming.await.expect("accepted WebRTC carrier connection");
assert_eq!(connection.alpn(), TEST_ALPN);
let (mut send, mut recv) = connection.accept_bi().await.unwrap();
assert_eq!(recv.read_to_end(1024).await.unwrap(), b"iroh-over-webrtc");
send.write_all(b"bilateral-iroh-response").await.unwrap();
send.finish().unwrap();
let _ = connection.closed().await;
})
};
let address_b =
EndpointAddr::new(endpoint_id_b).with_addrs([TransportAddr::Custom(remote_addr_b)]);
let connection = tokio::time::timeout(
std::time::Duration::from_secs(15),
endpoint_a.connect(address_b, TEST_ALPN),
)
.await
.expect("WebRTC-backed Iroh connect timed out")
.expect("WebRTC-backed Iroh connect failed");
assert!(connection.paths().iter().any(|path| {
path.is_selected()
&& matches!(
path.remote_addr(),
TransportAddr::Custom(address)
if address.id() == EXPERIMENTAL_WEBRTC_TRANSPORT_ID
)
}));
let (mut send, mut recv) = connection.open_bi().await.unwrap();
send.write_all(b"iroh-over-webrtc").await.unwrap();
send.finish().unwrap();
assert_eq!(
recv.read_to_end(1024).await.unwrap(),
b"bilateral-iroh-response"
);
connection.close(0u8.into(), b"test-complete");
server.await.unwrap();
crate::native_carrier_protocol_test::assert_docs_blobs_gossip_over_custom_transport(
endpoint_a.clone(),
endpoint_b.clone(),
)
.await
.unwrap();
carrier_a
.send_terminal(crate::lifecycle_reason::REASON_SESSION_TOKEN_REVOKED)
.await
.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
if carrier_b.take_terminal_reason()
== Some(crate::lifecycle_reason::REASON_SESSION_TOKEN_REVOKED)
{
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
})
.await
.expect("terminal WebRTC marker should be acknowledged and observed");
endpoint_a.close().await;
endpoint_b.close().await;
drop(carrier_a);
drop(carrier_b);
channel_a.close();
channel_b.close();
relay_a_to_b.abort();
relay_b_to_a.abort();
}
}