#![cfg(all(not(target_arch = "wasm32"), feature = "transport-moq"))]
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicU8, Ordering};
use std::sync::Arc;
use std::time::Duration;
use anyhow::{anyhow, bail, ensure, Context, Result};
use bytes::Bytes;
use tokio::sync::{Mutex, RwLock};
use super::{NativeMoQMessageHandler, NativeMoQState};
const MOQ_DRAFT_07_VERSION: u64 = 0xff00_0007;
const MSG_CLIENT_SETUP: u64 = 0x40;
const MSG_SERVER_SETUP: u64 = 0x41;
const MSG_SUBSCRIBE: u64 = 0x03;
const MSG_SUBSCRIBE_OK: u64 = 0x04;
const MSG_ANNOUNCE: u64 = 0x06;
const MSG_ANNOUNCE_OK: u64 = 0x07;
const MSG_ANNOUNCE_ERROR: u64 = 0x08;
const MSG_UNSUBSCRIBE: u64 = 0x0a;
const MSG_MAX_SUBSCRIBE_ID: u64 = 0x15;
const STREAM_HEADER_SUBGROUP: u64 = 0x04;
const SETUP_ROLE_PARAMETER_KEY: u64 = 0x00;
const SETUP_ROLE_BOTH: u64 = 3;
const MOQ_SETUP_TIMEOUT: Duration = Duration::from_secs(10);
const MOQ_MAX_CONTROL_FRAME_BYTES: usize = 64 * 1024;
const MOQ_MAX_OBJECT_PAYLOAD_BYTES: usize = 4 * 1024 * 1024;
#[derive(Debug, Clone)]
struct SubscribedTrack {
subscribe_id: u64,
track_alias: u64,
next_group_id: u64,
}
#[derive(Clone)]
pub struct NativeMoQSession {
local_node_id: String,
remote_node_id: String,
relay_url: String,
state: Arc<AtomicU8>,
started: Arc<AtomicBool>,
connection: Arc<Mutex<Option<wtransport::Connection>>>,
control_send: Arc<Mutex<Option<wtransport::SendStream>>>,
subscribed_tracks: Arc<Mutex<HashMap<String, SubscribedTrack>>>,
object_counter: Arc<AtomicU64>,
message_handler: Arc<RwLock<Option<NativeMoQMessageHandler>>>,
pending_messages: Arc<Mutex<Vec<Bytes>>>,
}
impl NativeMoQSession {
pub async fn new(
local_node_id: &str,
remote_node_id: &str,
config: &crate::client::MoQConfig,
) -> Result<Self> {
let relay_url = config.relay_url.trim();
ensure!(!relay_url.is_empty(), "moq relay_url is required");
ensure!(
relay_url.starts_with("https://"),
"moq relay_url must use https/WebTransport"
);
Ok(Self {
local_node_id: local_node_id.to_string(),
remote_node_id: remote_node_id.to_string(),
relay_url: relay_url.to_string(),
state: Arc::new(AtomicU8::new(0)),
started: Arc::new(AtomicBool::new(false)),
connection: Arc::new(Mutex::new(None)),
control_send: Arc::new(Mutex::new(None)),
subscribed_tracks: Arc::new(Mutex::new(HashMap::new())),
object_counter: Arc::new(AtomicU64::new(0)),
message_handler: Arc::new(RwLock::new(None)),
pending_messages: Arc::new(Mutex::new(Vec::new())),
})
}
pub fn state(&self) -> NativeMoQState {
match self.state.load(Ordering::SeqCst) {
1 => NativeMoQState::Connecting,
2 => NativeMoQState::Connected,
3 => NativeMoQState::Failed,
4 => NativeMoQState::Closed,
_ => NativeMoQState::Idle,
}
}
#[cfg(test)]
pub(crate) fn force_state_for_test(&self, state: NativeMoQState) {
let value = match state {
NativeMoQState::Idle => 0,
NativeMoQState::Connecting => 1,
NativeMoQState::Connected => 2,
NativeMoQState::Failed => 3,
NativeMoQState::Closed => 4,
};
self.state.store(value, Ordering::SeqCst);
}
pub async fn start(&self) -> Result<()> {
if self.started.swap(true, Ordering::SeqCst) {
return Ok(());
}
self.state.store(1, Ordering::SeqCst);
if self.relay_url.contains(".invalid") {
self.state.store(3, Ordering::SeqCst);
anyhow::bail!("moq relay is unreachable: {}", self.relay_url);
}
let result = tokio::time::timeout(MOQ_SETUP_TIMEOUT, async {
let client_config = if relay_url_uses_loopback_host(&self.relay_url) {
wtransport::ClientConfig::builder()
.with_bind_default()
.with_no_cert_validation()
.build()
} else {
wtransport::ClientConfig::builder()
.with_bind_default()
.with_native_certs()
.build()
};
let connection = wtransport::Endpoint::client(client_config)
.context("failed to create MoQ WebTransport endpoint")?
.connect(self.relay_url.as_str())
.await
.with_context(|| format!("failed to connect MoQ relay {}", self.relay_url))?;
let (mut control_send, mut control_recv) = connection
.open_bi()
.await
.context("failed to open MoQ control stream")?
.await
.context("failed to establish MoQ control stream")?;
let setup = encode_client_setup();
control_send
.write_all(&setup)
.await
.context("failed to send MoQ CLIENT_SETUP")?;
let (selected_version, role) = read_server_setup(&mut control_recv)
.await
.context("failed to read MoQ SERVER_SETUP")?;
ensure!(
selected_version == MOQ_DRAFT_07_VERSION,
"MoQ relay selected unsupported version 0x{selected_version:x}"
);
ensure!(
role == SETUP_ROLE_BOTH || role == 0,
"MoQ relay selected unsupported role {role}"
);
control_send
.write_all(&encode_announce(&self.local_node_id))
.await
.context("failed to send MoQ ANNOUNCE")?;
let pending_control_frames = read_announce_ok(&mut control_recv, &self.local_node_id)
.await
.context("failed to read MoQ ANNOUNCE_OK")?;
write_peer_data_subscriptions(
&mut control_send,
&self.remote_node_id,
&self.local_node_id,
)
.await?;
Ok::<
(
wtransport::Connection,
wtransport::SendStream,
wtransport::RecvStream,
Vec<(u64, Vec<u8>)>,
),
anyhow::Error,
>((
connection,
control_send,
control_recv,
pending_control_frames,
))
})
.await
.map_err(|_| {
anyhow!(
"MoQ setup timed out after {} ms",
MOQ_SETUP_TIMEOUT.as_millis()
)
});
match result {
Ok(Ok((connection, control_send, control_recv, pending_control_frames))) => {
*self.control_send.lock().await = Some(control_send);
*self.connection.lock().await = Some(connection);
self.state.store(2, Ordering::SeqCst);
for frame in pending_control_frames {
handle_control_frame(
frame,
&self.local_node_id,
&self.control_send,
&self.subscribed_tracks,
)
.await
.context("failed to handle buffered MoQ control frame")?;
}
self.spawn_control_loop(control_recv);
self.spawn_unistream_loop();
self.spawn_subscription_refresh_loop();
Ok(())
}
Ok(Err(error)) => {
self.state.store(3, Ordering::SeqCst);
Err(error)
}
Err(error) => {
self.state.store(3, Ordering::SeqCst);
Err(error)
}
}
}
pub async fn send(&self, data: &[u8]) -> Result<()> {
anyhow::ensure!(
self.state() == NativeMoQState::Connected,
"native moq session is not started"
);
ensure!(
data.len() <= MOQ_MAX_OBJECT_PAYLOAD_BYTES,
"native moq object payload too large: {}",
data.len()
);
let track_name = self.peer_data_track_name();
let mut subscribed_tracks = self.subscribed_tracks.lock().await;
let track = subscribed_tracks.get_mut(&track_name).cloned();
let Some(mut track) = track else {
bail!("native moq track is not subscribed yet; caller must fall back to iroh");
};
track.next_group_id = track.next_group_id.saturating_add(1);
subscribed_tracks.insert(track_name.clone(), track.clone());
drop(subscribed_tracks);
let connection = {
let guard = self.connection.lock().await;
guard.clone()
};
let Some(connection) = connection else {
bail!("native moq connection is not available")
};
let mut stream = connection
.open_uni()
.await
.context("failed to open MoQ object stream")?
.await
.context("failed to establish MoQ object stream")?;
let object_id = self.object_counter.fetch_add(1, Ordering::SeqCst);
let object = encode_subgroup_object(
track.subscribe_id,
track.track_alias,
track.next_group_id,
object_id,
data,
);
stream
.write_all(&object)
.await
.context("failed to write MoQ object")?;
stream
.finish()
.await
.context("failed to finish MoQ object stream")?;
Ok(())
}
pub async fn is_peer_data_subscribed(&self) -> bool {
let track_name = self.peer_data_track_name();
self.subscribed_tracks
.lock()
.await
.contains_key(&track_name)
}
pub async fn wait_for_peer_data_subscription(&self, timeout: Duration) -> bool {
let track_name = self.peer_data_track_name();
tokio::time::timeout(timeout, async {
loop {
if self
.subscribed_tracks
.lock()
.await
.contains_key(&track_name)
{
return true;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
})
.await
.unwrap_or(false)
}
pub async fn close_gracefully(&self) {
let connection = self.connection.clone();
let control_send = self.control_send.clone();
let _ = control_send.lock().await.take();
if let Some(connection) = connection.lock().await.take() {
connection.close(wtransport::VarInt::from_u32(0), b"local-close");
}
self.state.store(4, Ordering::SeqCst);
}
pub fn close(&self) {
let session = self.clone();
tokio::spawn(async move {
session.close_gracefully().await;
});
self.state.store(4, Ordering::SeqCst);
}
pub async fn set_message_handler(&self, handler: NativeMoQMessageHandler) {
{
let mut guard = self.message_handler.write().await;
*guard = Some(handler.clone());
}
let pending = {
let mut guard = self.pending_messages.lock().await;
std::mem::take(&mut *guard)
};
if !pending.is_empty() {
println!(
"[NativeMoQ] flushing {} buffered object message(s) after handler registration",
pending.len()
);
}
for message in pending {
handler(message).await;
}
}
fn peer_data_track_name(&self) -> String {
format!("data / {} ", self.remote_node_id)
}
fn spawn_subscription_refresh_loop(&self) {
let control_send = self.control_send.clone();
let state = self.state.clone();
let remote_node_id = self.remote_node_id.clone();
let local_node_id = self.local_node_id.clone();
tokio::spawn(async move {
for _ in 0..20 {
tokio::time::sleep(Duration::from_millis(500)).await;
if state.load(Ordering::SeqCst) != 2 {
return;
}
let mut guard = control_send.lock().await;
let Some(send) = guard.as_mut() else {
return;
};
if let Err(error) =
write_peer_data_subscriptions(send, &remote_node_id, &local_node_id).await
{
eprintln!("[NativeMoQ] subscription refresh failed: {error}");
return;
}
}
});
}
fn spawn_control_loop(&self, mut control_recv: wtransport::RecvStream) {
let local_node_id = self.local_node_id.clone();
let state = self.state.clone();
let control_send = self.control_send.clone();
let subscribed_tracks = self.subscribed_tracks.clone();
tokio::spawn(async move {
let mut buf = Vec::new();
let mut chunk = [0u8; 2048];
loop {
let read = match control_recv.read(&mut chunk).await {
Ok(Some(read)) => read,
Ok(None) => {
state.store(4, Ordering::SeqCst);
return;
}
Err(error) => {
eprintln!("[NativeMoQ] control read failed: {error}");
state.store(3, Ordering::SeqCst);
return;
}
};
buf.extend_from_slice(&chunk[..read]);
loop {
let frame = match try_decode_control_frame(&mut buf) {
Ok(Some(frame)) => frame,
Ok(None) => break,
Err(error) => {
eprintln!("[NativeMoQ] control parse failed: {error}");
state.store(3, Ordering::SeqCst);
return;
}
};
if let Err(error) = handle_control_frame(
frame,
&local_node_id,
&control_send,
&subscribed_tracks,
)
.await
{
eprintln!("[NativeMoQ] control handling failed: {error}");
state.store(3, Ordering::SeqCst);
return;
}
}
}
});
}
fn spawn_unistream_loop(&self) {
let connection = self.connection.clone();
let state = self.state.clone();
let message_handler = self.message_handler.clone();
let pending_messages = self.pending_messages.clone();
tokio::spawn(async move {
let connection = {
let guard = connection.lock().await;
guard.clone()
};
let Some(connection) = connection else {
return;
};
loop {
let mut stream = match connection.accept_uni().await {
Ok(stream) => stream,
Err(error) => {
eprintln!("[NativeMoQ] accept_uni failed: {error}");
if state.load(Ordering::SeqCst) != 4 {
state.store(3, Ordering::SeqCst);
}
return;
}
};
match read_moq_object_stream(&mut stream).await {
Ok(objects) => {
for object in objects {
dispatch_or_buffer_message(
&message_handler,
&pending_messages,
Bytes::from(object),
)
.await;
}
}
Err(error) => {
eprintln!("[NativeMoQ] object stream parse failed: {error}");
}
}
}
});
}
}
fn relay_url_uses_loopback_host(relay_url: &str) -> bool {
relay_url.starts_with("https://localhost")
|| relay_url.starts_with("https://127.")
|| relay_url.starts_with("https://[::1]")
}
fn write_varint(out: &mut Vec<u8>, value: u64) {
if value < 64 {
out.push(value as u8);
} else if value < 16_384 {
out.push(((value >> 8) as u8) | 0x40);
out.push(value as u8);
} else if value < 1_073_741_824 {
out.push(((value >> 24) as u8) | 0x80);
out.push((value >> 16) as u8);
out.push((value >> 8) as u8);
out.push(value as u8);
} else {
out.push(((value >> 56) as u8) | 0xc0);
out.push((value >> 48) as u8);
out.push((value >> 40) as u8);
out.push((value >> 32) as u8);
out.push((value >> 24) as u8);
out.push((value >> 16) as u8);
out.push((value >> 8) as u8);
out.push(value as u8);
}
}
fn read_varint(buf: &[u8], offset: &mut usize) -> Result<u64> {
ensure!(*offset < buf.len(), "unexpected EOF reading MoQ varint");
let first = buf[*offset];
let tag = first >> 6;
let len = 1usize << tag;
ensure!(
*offset + len <= buf.len(),
"unexpected EOF reading MoQ varint"
);
let mut value = (first & 0x3f) as u64;
for byte in &buf[*offset + 1..*offset + len] {
value = (value << 8) | u64::from(*byte);
}
*offset += len;
Ok(value)
}
fn encode_frame(message_type: u64, payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(16 + payload.len());
write_varint(&mut out, message_type);
write_varint(&mut out, payload.len() as u64);
out.extend_from_slice(payload);
out
}
fn encode_client_setup() -> Vec<u8> {
let mut payload = Vec::new();
write_varint(&mut payload, 1);
write_varint(&mut payload, MOQ_DRAFT_07_VERSION);
write_varint(&mut payload, 1);
write_varint(&mut payload, SETUP_ROLE_PARAMETER_KEY);
write_varint(&mut payload, 1);
write_varint(&mut payload, SETUP_ROLE_BOTH);
encode_frame(MSG_CLIENT_SETUP, &payload)
}
fn write_bytes(out: &mut Vec<u8>, bytes: &[u8]) {
write_varint(out, bytes.len() as u64);
out.extend_from_slice(bytes);
}
fn write_namespace_tuple(out: &mut Vec<u8>, namespace: &str) {
write_varint(out, 1);
write_bytes(out, namespace.as_bytes());
}
fn read_bytes<'a>(buf: &'a [u8], offset: &mut usize) -> Result<&'a [u8]> {
let len = read_varint(buf, offset)? as usize;
ensure!(
*offset + len <= buf.len(),
"unexpected EOF reading MoQ bytes"
);
let bytes = &buf[*offset..*offset + len];
*offset += len;
Ok(bytes)
}
fn read_namespace_tuple(buf: &[u8], offset: &mut usize) -> Result<String> {
let part_count = read_varint(buf, offset)? as usize;
let mut parts = Vec::with_capacity(part_count);
for _ in 0..part_count {
let bytes = read_bytes(buf, offset)?;
parts.push(
std::str::from_utf8(bytes)
.context("MoQ namespace is not valid UTF-8")?
.to_string(),
);
}
Ok(parts.join("/"))
}
fn encode_announce(namespace: &str) -> Vec<u8> {
let mut payload = Vec::new();
write_namespace_tuple(&mut payload, namespace);
write_varint(&mut payload, 0);
encode_frame(MSG_ANNOUNCE, &payload)
}
fn encode_subscribe(
subscribe_id: u64,
track_alias: u64,
namespace: &str,
track_name: &str,
) -> Vec<u8> {
let mut payload = Vec::new();
write_varint(&mut payload, subscribe_id);
write_varint(&mut payload, track_alias);
write_namespace_tuple(&mut payload, namespace);
write_bytes(&mut payload, track_name.as_bytes());
payload.push(128);
payload.push(0x01);
write_varint(&mut payload, 0x01);
write_varint(&mut payload, 0);
encode_frame(MSG_SUBSCRIBE, &payload)
}
#[cfg(test)]
fn encode_announce_ok_for_tests(namespace: &str) -> Vec<u8> {
let mut payload = Vec::new();
write_namespace_tuple(&mut payload, namespace);
encode_frame(MSG_ANNOUNCE_OK, &payload)
}
#[cfg(test)]
fn encode_subscribe_for_tests(
subscribe_id: u64,
track_alias: u64,
namespace: &str,
track_name: &str,
) -> Vec<u8> {
encode_subscribe(subscribe_id, track_alias, namespace, track_name)
}
fn encode_subscribe_ok(subscribe_id: u64) -> Vec<u8> {
let mut payload = Vec::new();
write_varint(&mut payload, subscribe_id);
write_varint(&mut payload, 0);
payload.push(0x01);
payload.push(0);
write_varint(&mut payload, 0);
encode_frame(MSG_SUBSCRIBE_OK, &payload)
}
fn encode_subgroup_object(
subscribe_id: u64,
track_alias: u64,
group_id: u64,
object_id: u64,
data: &[u8],
) -> Vec<u8> {
let mut out = Vec::with_capacity(32 + data.len());
write_varint(&mut out, STREAM_HEADER_SUBGROUP);
write_varint(&mut out, subscribe_id);
write_varint(&mut out, track_alias);
write_varint(&mut out, group_id);
write_varint(&mut out, 0);
out.push(0);
write_varint(&mut out, object_id);
write_varint(&mut out, data.len() as u64);
out.extend_from_slice(data);
out
}
#[cfg(test)]
fn encode_server_setup_for_tests(selected_version: u64, role: u64) -> Vec<u8> {
let mut payload = Vec::new();
write_varint(&mut payload, selected_version);
write_varint(&mut payload, 1);
write_varint(&mut payload, SETUP_ROLE_PARAMETER_KEY);
write_varint(&mut payload, 1);
write_varint(&mut payload, role);
encode_frame(MSG_SERVER_SETUP, &payload)
}
fn try_decode_control_frame(buf: &mut Vec<u8>) -> Result<Option<(u64, Vec<u8>)>> {
let mut offset = 0usize;
let message_type = match read_varint(buf, &mut offset) {
Ok(value) => value,
Err(error)
if error.to_string().contains("incomplete") || error.to_string().contains("EOF") =>
{
return Ok(None);
}
Err(error) => return Err(error),
};
let payload_len = match read_varint(buf, &mut offset) {
Ok(value) => value as usize,
Err(error)
if error.to_string().contains("incomplete") || error.to_string().contains("EOF") =>
{
return Ok(None);
}
Err(error) => return Err(error),
};
ensure!(
payload_len <= MOQ_MAX_CONTROL_FRAME_BYTES,
"MoQ control frame too large: {payload_len}"
);
if offset + payload_len > buf.len() {
return Ok(None);
}
let payload = buf[offset..offset + payload_len].to_vec();
buf.drain(..offset + payload_len);
Ok(Some((message_type, payload)))
}
fn parse_server_setup_frame(frame: &[u8]) -> Result<(u64, u64)> {
let mut offset = 0usize;
let message_type = read_varint(frame, &mut offset)?;
ensure!(
message_type == MSG_SERVER_SETUP,
"expected MoQ SERVER_SETUP, got 0x{message_type:x}"
);
let payload_len = read_varint(frame, &mut offset)? as usize;
ensure!(
payload_len <= MOQ_MAX_CONTROL_FRAME_BYTES,
"MoQ control frame too large: {payload_len}"
);
ensure!(
offset + payload_len <= frame.len(),
"incomplete MoQ SERVER_SETUP payload"
);
let payload = &frame[offset..offset + payload_len];
let mut payload_offset = 0usize;
let selected_version = read_varint(payload, &mut payload_offset)?;
let param_count = read_varint(payload, &mut payload_offset)?;
let mut role = 0u64;
for _ in 0..param_count {
let key = read_varint(payload, &mut payload_offset)?;
let len = read_varint(payload, &mut payload_offset)? as usize;
ensure!(
payload_offset + len <= payload.len(),
"incomplete MoQ SERVER_SETUP parameter"
);
if key == SETUP_ROLE_PARAMETER_KEY {
let param_end = payload_offset + len;
let mut param_offset = payload_offset;
role = read_varint(payload, &mut param_offset)?;
ensure!(
param_offset <= param_end,
"MoQ SERVER_SETUP role parameter exceeded declared length"
);
}
payload_offset += len;
}
Ok((selected_version, role))
}
fn parse_announce_ok_payload(payload: &[u8]) -> Result<String> {
let mut offset = 0usize;
read_namespace_tuple(payload, &mut offset)
}
fn parse_announce_error_payload(payload: &[u8]) -> Result<(String, u64, String)> {
let mut offset = 0usize;
let namespace = read_namespace_tuple(payload, &mut offset)?;
let code = read_varint(payload, &mut offset)?;
let reason = std::str::from_utf8(read_bytes(payload, &mut offset)?)
.context("MoQ ANNOUNCE_ERROR reason is not valid UTF-8")?
.to_string();
Ok((namespace, code, reason))
}
#[derive(Debug)]
struct SubscribeMessage {
subscribe_id: u64,
track_alias: u64,
namespace: String,
track_name: String,
}
fn parse_subscribe_payload(payload: &[u8]) -> Result<SubscribeMessage> {
let mut offset = 0usize;
let subscribe_id = read_varint(payload, &mut offset)?;
let track_alias = read_varint(payload, &mut offset)?;
let namespace = read_namespace_tuple(payload, &mut offset)?;
let track_name = std::str::from_utf8(read_bytes(payload, &mut offset)?)
.context("MoQ SUBSCRIBE track name is not valid UTF-8")?
.to_string();
ensure!(
offset + 2 <= payload.len(),
"incomplete MoQ SUBSCRIBE priority/order"
);
offset += 2;
let filter_type = read_varint(payload, &mut offset)?;
match filter_type {
0x03 => {
let _ = read_varint(payload, &mut offset)?;
let _ = read_varint(payload, &mut offset)?;
}
0x04 => {
let _ = read_varint(payload, &mut offset)?;
let _ = read_varint(payload, &mut offset)?;
let _ = read_varint(payload, &mut offset)?;
let _ = read_varint(payload, &mut offset)?;
}
_ => {}
}
let param_count = read_varint(payload, &mut offset)?;
for _ in 0..param_count {
let _key = read_varint(payload, &mut offset)?;
let len = read_varint(payload, &mut offset)? as usize;
ensure!(
offset + len <= payload.len(),
"incomplete MoQ SUBSCRIBE parameter"
);
offset += len;
}
Ok(SubscribeMessage {
subscribe_id,
track_alias,
namespace,
track_name,
})
}
async fn write_control_frame(
control_send: &Arc<Mutex<Option<wtransport::SendStream>>>,
frame: &[u8],
) -> Result<()> {
let mut guard = control_send.lock().await;
let Some(send) = guard.as_mut() else {
bail!("MoQ control stream is not available");
};
send.write_all(frame)
.await
.context("failed to write MoQ control frame")
}
async fn write_peer_data_subscriptions(
control_send: &mut wtransport::SendStream,
remote_node_id: &str,
local_node_id: &str,
) -> Result<()> {
control_send
.write_all(&encode_subscribe(1, 0, remote_node_id, "data/broadcast"))
.await
.context("failed to send MoQ broadcast SUBSCRIBE")?;
control_send
.write_all(&encode_subscribe(
2,
1,
remote_node_id,
&format!("data / {} ", local_node_id),
))
.await
.context("failed to send MoQ direct SUBSCRIBE")?;
Ok(())
}
async fn handle_control_frame(
frame: (u64, Vec<u8>),
local_node_id: &str,
control_send: &Arc<Mutex<Option<wtransport::SendStream>>>,
subscribed_tracks: &Arc<Mutex<HashMap<String, SubscribedTrack>>>,
) -> Result<()> {
match frame.0 {
MSG_SUBSCRIBE => {
let subscribe = parse_subscribe_payload(&frame.1)?;
if subscribe.namespace == local_node_id {
subscribed_tracks.lock().await.insert(
subscribe.track_name,
SubscribedTrack {
subscribe_id: subscribe.subscribe_id,
track_alias: subscribe.track_alias,
next_group_id: 0,
},
);
}
write_control_frame(control_send, &encode_subscribe_ok(subscribe.subscribe_id)).await?;
}
MSG_UNSUBSCRIBE => {
let mut offset = 0usize;
let subscribe_id = read_varint(&frame.1, &mut offset)?;
subscribed_tracks
.lock()
.await
.retain(|_, track| track.subscribe_id != subscribe_id);
}
MSG_MAX_SUBSCRIBE_ID | MSG_ANNOUNCE_OK => {}
MSG_ANNOUNCE_ERROR => {
let (namespace, code, reason) = parse_announce_error_payload(&frame.1)?;
bail!("MoQ announce failed for {namespace}: code={code} reason={reason}");
}
_ => {}
}
Ok(())
}
async fn dispatch_or_buffer_message(
handler: &Arc<RwLock<Option<NativeMoQMessageHandler>>>,
pending_messages: &Arc<Mutex<Vec<Bytes>>>,
message: Bytes,
) {
let callback = {
let guard = handler.read().await;
guard.clone()
};
if let Some(callback) = callback {
callback(message).await;
return;
}
pending_messages.lock().await.push(message);
}
async fn read_moq_object_stream(recv: &mut wtransport::RecvStream) -> Result<Vec<Vec<u8>>> {
let mut buf = Vec::new();
let mut chunk = [0u8; 4096];
loop {
match recv
.read(&mut chunk)
.await
.context("failed to read MoQ object stream")?
{
Some(read) => buf.extend_from_slice(&chunk[..read]),
None => break,
}
ensure!(
buf.len() <= MOQ_MAX_OBJECT_PAYLOAD_BYTES + 128,
"MoQ object stream exceeded maximum payload bound"
);
}
let mut offset = 0usize;
let stream_type = read_varint(&buf, &mut offset)?;
ensure!(
stream_type == STREAM_HEADER_SUBGROUP,
"unexpected MoQ object stream type: {stream_type}"
);
let _subscribe_id = read_varint(&buf, &mut offset)?;
let _track_alias = read_varint(&buf, &mut offset)?;
let group_id = read_varint(&buf, &mut offset)?;
let _subgroup_id = read_varint(&buf, &mut offset)?;
ensure!(offset < buf.len(), "incomplete MoQ object stream priority");
let _priority = buf[offset];
offset += 1;
let mut objects = Vec::new();
while offset < buf.len() {
let object_id = read_varint(&buf, &mut offset)?;
let len = read_varint(&buf, &mut offset)? as usize;
ensure!(
len <= MOQ_MAX_OBJECT_PAYLOAD_BYTES,
"MoQ object payload too large: {len}"
);
ensure!(offset + len <= buf.len(), "incomplete MoQ object payload");
objects.push(buf[offset..offset + len].to_vec());
offset += len;
let _ = (group_id, object_id);
}
Ok(objects)
}
async fn read_server_setup(recv: &mut wtransport::RecvStream) -> Result<(u64, u64)> {
let mut buf = Vec::new();
let mut chunk = [0u8; 2048];
loop {
let read = recv
.read(&mut chunk)
.await
.context("failed to read MoQ control stream")?;
let Some(read) = read else {
bail!("MoQ control stream closed before SERVER_SETUP");
};
buf.extend_from_slice(&chunk[..read]);
if buf.len() > MOQ_MAX_CONTROL_FRAME_BYTES {
bail!("MoQ control stream exceeded {MOQ_MAX_CONTROL_FRAME_BYTES} bytes before SERVER_SETUP");
}
match parse_server_setup_frame(&buf) {
Ok(parsed) => return Ok(parsed),
Err(error)
if error.to_string().contains("incomplete")
|| error.to_string().contains("EOF") =>
{
continue;
}
Err(error) => return Err(error),
}
}
}
async fn read_announce_ok(
recv: &mut wtransport::RecvStream,
local_node_id: &str,
) -> Result<Vec<(u64, Vec<u8>)>> {
let mut buf = Vec::new();
let mut pending_frames = Vec::new();
let mut chunk = [0u8; 2048];
loop {
let read = recv
.read(&mut chunk)
.await
.context("failed to read MoQ control stream")?;
let Some(read) = read else {
bail!("MoQ control stream closed before ANNOUNCE_OK");
};
buf.extend_from_slice(&chunk[..read]);
if buf.len() > MOQ_MAX_CONTROL_FRAME_BYTES {
bail!("MoQ control stream exceeded {MOQ_MAX_CONTROL_FRAME_BYTES} bytes before ANNOUNCE_OK");
}
while let Some((message_type, payload)) = try_decode_control_frame(&mut buf)? {
match message_type {
MSG_ANNOUNCE_OK => {
let namespace = parse_announce_ok_payload(&payload)?;
ensure!(
namespace == local_node_id,
"MoQ ANNOUNCE_OK namespace mismatch: expected {local_node_id}, got {namespace}"
);
while let Some(frame) = try_decode_control_frame(&mut buf)? {
pending_frames.push(frame);
}
return Ok(pending_frames);
}
MSG_ANNOUNCE_ERROR => {
let (namespace, code, reason) = parse_announce_error_payload(&payload)?;
bail!("MoQ announce failed for {namespace}: code={code} reason={reason}");
}
_ => pending_frames.push((message_type, payload)),
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::client::MoQConfig;
use std::collections::HashMap;
use std::time::Duration;
use wtransport::{Endpoint, Identity, ServerConfig};
fn make_config(relay_url: &str) -> MoQConfig {
MoQConfig {
relay_url: relay_url.to_string(),
}
}
#[derive(Default)]
struct LoopbackRelayState {
peers: HashMap<String, Arc<LoopbackRelayPeer>>,
routes: HashMap<(String, u64, u64), wtransport::Connection>,
pending_subscribes: HashMap<String, Vec<Vec<u8>>>,
}
struct LoopbackRelayPeer {
control_send: Arc<Mutex<wtransport::SendStream>>,
}
fn object_route_key(publisher_node_id: &str, object: &[u8]) -> Option<(String, u64, u64)> {
let mut offset = 0usize;
if read_varint(object, &mut offset).ok()? != STREAM_HEADER_SUBGROUP {
return None;
}
let subscribe_id = read_varint(object, &mut offset).ok()?;
let track_alias = read_varint(object, &mut offset).ok()?;
Some((publisher_node_id.to_string(), subscribe_id, track_alias))
}
async fn read_entire_uni_stream(mut stream: wtransport::RecvStream) -> Vec<u8> {
let mut buf = Vec::new();
let mut chunk = [0u8; 2048];
loop {
match stream.read(&mut chunk).await {
Ok(Some(read)) => buf.extend_from_slice(&chunk[..read]),
Ok(None) | Err(_) => break,
}
}
buf
}
async fn forward_pending_subscribes(
state: &Arc<Mutex<LoopbackRelayState>>,
publisher_node_id: &str,
) {
let (peer, pending) = {
let mut guard = state.lock().await;
(
guard.peers.get(publisher_node_id).cloned(),
guard
.pending_subscribes
.remove(publisher_node_id)
.unwrap_or_default(),
)
};
let Some(peer) = peer else {
return;
};
for subscribe in pending {
let _ = peer.control_send.lock().await.write_all(&subscribe).await;
}
}
async fn process_loopback_relay_subscribes(
control_buf: &mut Vec<u8>,
subscriber_connection: &wtransport::Connection,
state: &Arc<Mutex<LoopbackRelayState>>,
) {
while let Ok(Some((message_type, payload))) = try_decode_control_frame(control_buf) {
if message_type != MSG_SUBSCRIBE {
continue;
}
let Ok(subscribe) = parse_subscribe_payload(&payload) else {
continue;
};
let encoded = encode_frame(MSG_SUBSCRIBE, &payload);
let publisher = {
let mut guard = state.lock().await;
guard.routes.insert(
(
subscribe.namespace.clone(),
subscribe.subscribe_id,
subscribe.track_alias,
),
subscriber_connection.clone(),
);
match guard.peers.get(&subscribe.namespace).cloned() {
Some(peer) => Some(peer),
None => {
guard
.pending_subscribes
.entry(subscribe.namespace.clone())
.or_default()
.push(encoded.clone());
None
}
}
};
if let Some(publisher) = publisher {
let _ = publisher
.control_send
.lock()
.await
.write_all(&encoded)
.await;
}
}
}
async fn handle_loopback_relay_connection(
connection: wtransport::Connection,
mut control_send: wtransport::SendStream,
mut control_recv: wtransport::RecvStream,
state: Arc<Mutex<LoopbackRelayState>>,
) {
let mut control_buf = Vec::new();
let mut chunk = [0u8; 2048];
let mut announced_node_id = None;
while announced_node_id.is_none() {
let read = match control_recv.read(&mut chunk).await {
Ok(Some(read)) => read,
_ => return,
};
control_buf.extend_from_slice(&chunk[..read]);
while let Ok(Some((message_type, payload))) = try_decode_control_frame(&mut control_buf)
{
match message_type {
MSG_CLIENT_SETUP => {
let _ = control_send
.write_all(&encode_server_setup_for_tests(
MOQ_DRAFT_07_VERSION,
SETUP_ROLE_BOTH,
))
.await;
}
MSG_ANNOUNCE => {
let mut offset = 0usize;
let Ok(node_id) = read_namespace_tuple(&payload, &mut offset) else {
return;
};
let _ = control_send
.write_all(&encode_announce_ok_for_tests(&node_id))
.await;
announced_node_id = Some(node_id);
break;
}
_ => {}
}
}
}
let Some(node_id) = announced_node_id else {
return;
};
let control_send = Arc::new(Mutex::new(control_send));
{
let mut guard = state.lock().await;
guard.peers.insert(
node_id.clone(),
Arc::new(LoopbackRelayPeer {
control_send: control_send.clone(),
}),
);
}
forward_pending_subscribes(&state, &node_id).await;
process_loopback_relay_subscribes(&mut control_buf, &connection, &state).await;
{
let state = state.clone();
let connection = connection.clone();
let publisher_node_id = node_id.clone();
tokio::spawn(async move {
loop {
let stream = match connection.accept_uni().await {
Ok(stream) => stream,
Err(_) => break,
};
let object = read_entire_uni_stream(stream).await;
let Some(route_key) = object_route_key(&publisher_node_id, &object) else {
continue;
};
let subscriber = {
let guard = state.lock().await;
guard.routes.get(&route_key).cloned()
};
let Some(subscriber) = subscriber else {
continue;
};
let Ok(opening) = subscriber.open_uni().await else {
continue;
};
let Ok(mut outbound) = opening.await else {
continue;
};
if outbound.write_all(&object).await.is_ok() {
let _ = outbound.finish().await;
}
}
});
}
loop {
let read = match control_recv.read(&mut chunk).await {
Ok(Some(read)) => read,
_ => break,
};
control_buf.extend_from_slice(&chunk[..read]);
process_loopback_relay_subscribes(&mut control_buf, &connection, &state).await;
}
}
async fn wait_for_direct_subscription(session: &NativeMoQSession) {
let track_name = session.peer_data_track_name();
tokio::time::timeout(Duration::from_secs(2), async {
loop {
if session
.subscribed_tracks
.lock()
.await
.contains_key(&track_name)
{
return;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
})
.await
.expect("direct MoQ subscription should be forwarded to publisher");
}
#[test]
fn client_setup_wire_format_matches_draft_shape() {
let setup = encode_client_setup();
let mut offset = 0usize;
assert_eq!(read_varint(&setup, &mut offset).unwrap(), MSG_CLIENT_SETUP);
let payload_len = read_varint(&setup, &mut offset).unwrap() as usize;
let payload = &setup[offset..offset + payload_len];
let mut payload_offset = 0usize;
assert_eq!(read_varint(payload, &mut payload_offset).unwrap(), 1);
assert_eq!(
read_varint(payload, &mut payload_offset).unwrap(),
MOQ_DRAFT_07_VERSION
);
assert_eq!(read_varint(payload, &mut payload_offset).unwrap(), 1);
assert_eq!(
read_varint(payload, &mut payload_offset).unwrap(),
SETUP_ROLE_PARAMETER_KEY
);
assert_eq!(read_varint(payload, &mut payload_offset).unwrap(), 1);
assert_eq!(
read_varint(payload, &mut payload_offset).unwrap(),
SETUP_ROLE_BOTH
);
assert_eq!(payload_offset, payload.len());
}
#[test]
fn server_setup_parser_reads_selected_version_and_role() {
let frame = encode_server_setup_for_tests(MOQ_DRAFT_07_VERSION, SETUP_ROLE_BOTH);
assert_eq!(
parse_server_setup_frame(&frame).unwrap(),
(MOQ_DRAFT_07_VERSION, SETUP_ROLE_BOTH)
);
}
#[test]
fn valid_local_relay_reaches_connected_after_setup() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let server_config = ServerConfig::builder()
.with_bind_default(0)
.with_identity(Identity::self_signed(["localhost"]).unwrap())
.keep_alive_interval(Some(Duration::from_secs(3)))
.build();
let server = Endpoint::server(server_config).unwrap();
let addr = server.local_addr().unwrap();
let server_task = tokio::spawn(async move {
let incoming = server.accept().await;
let request = incoming.await.unwrap();
let connection = request.accept().await.unwrap();
let (mut send, mut recv) = connection.accept_bi().await.unwrap();
let mut control_buf = Vec::new();
let mut chunk = [0u8; 256];
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
let setup = try_decode_control_frame(&mut control_buf).unwrap().unwrap();
assert_eq!(setup.0, MSG_CLIENT_SETUP);
send.write_all(&encode_server_setup_for_tests(
MOQ_DRAFT_07_VERSION,
SETUP_ROLE_BOTH,
))
.await
.unwrap();
loop {
if let Some(frame) = try_decode_control_frame(&mut control_buf).unwrap() {
assert_eq!(frame.0, MSG_ANNOUNCE);
break;
}
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
}
send.write_all(&encode_announce_ok_for_tests("local"))
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(25)).await;
send.write_all(&encode_subscribe_for_tests(
7,
11,
"local",
"data / remote ",
))
.await
.unwrap();
let mut object_stream = connection.accept_uni().await.unwrap();
let mut object_buf = Vec::new();
let mut object_chunk = [0u8; 256];
loop {
match object_stream.read(&mut object_chunk).await.unwrap() {
Some(read) => object_buf.extend_from_slice(&object_chunk[..read]),
None => break,
}
}
let mut object_offset = 0usize;
assert_eq!(
read_varint(&object_buf, &mut object_offset).unwrap(),
STREAM_HEADER_SUBGROUP
);
assert_eq!(read_varint(&object_buf, &mut object_offset).unwrap(), 7);
assert_eq!(read_varint(&object_buf, &mut object_offset).unwrap(), 11);
assert_eq!(read_varint(&object_buf, &mut object_offset).unwrap(), 1);
assert_eq!(read_varint(&object_buf, &mut object_offset).unwrap(), 0);
assert_eq!(object_buf[object_offset], 0);
object_offset += 1;
assert_eq!(read_varint(&object_buf, &mut object_offset).unwrap(), 0);
let payload_len = read_varint(&object_buf, &mut object_offset).unwrap() as usize;
assert_eq!(
&object_buf[object_offset..object_offset + payload_len],
b"hello"
);
});
let session = NativeMoQSession::new(
"local",
"remote",
&make_config(&format!("https://localhost:{}/moq", addr.port())),
)
.await
.unwrap();
session.start().await.unwrap();
assert_eq!(session.state(), NativeMoQState::Connected);
tokio::time::sleep(Duration::from_millis(75)).await;
session.send(b"hello").await.unwrap();
server_task.await.unwrap();
});
}
#[test]
fn two_sessions_exchange_objects_through_loopback_relay() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let server_config = ServerConfig::builder()
.with_bind_default(0)
.with_identity(Identity::self_signed(["localhost"]).unwrap())
.keep_alive_interval(Some(Duration::from_secs(3)))
.build();
let server = Endpoint::server(server_config).unwrap();
let addr = server.local_addr().unwrap();
let relay_state = Arc::new(Mutex::new(LoopbackRelayState::default()));
let relay_task = {
let relay_state = relay_state.clone();
tokio::spawn(async move {
loop {
let incoming = server.accept().await;
let request = match incoming.await {
Ok(request) => request,
Err(_) => break,
};
let relay_state = relay_state.clone();
tokio::spawn(async move {
let Ok(connection) = request.accept().await else {
return;
};
let Ok((control_send, control_recv)) = connection.accept_bi().await
else {
return;
};
handle_loopback_relay_connection(
connection,
control_send,
control_recv,
relay_state,
)
.await;
});
}
})
};
let relay_url = format!("https://localhost:{}/moq", addr.port());
let node_a = NativeMoQSession::new("node-a", "node-b", &make_config(&relay_url))
.await
.unwrap();
let node_b = NativeMoQSession::new("node-b", "node-a", &make_config(&relay_url))
.await
.unwrap();
let (start_a, start_b) = tokio::join!(node_a.start(), node_b.start());
start_a.unwrap();
start_b.unwrap();
assert_eq!(node_a.state(), NativeMoQState::Connected);
assert_eq!(node_b.state(), NativeMoQState::Connected);
wait_for_direct_subscription(&node_a).await;
wait_for_direct_subscription(&node_b).await;
let (a_tx, a_rx) = tokio::sync::oneshot::channel::<Vec<u8>>();
let (b_tx, b_rx) = tokio::sync::oneshot::channel::<Vec<u8>>();
let a_tx = Arc::new(Mutex::new(Some(a_tx)));
let b_tx = Arc::new(Mutex::new(Some(b_tx)));
node_a
.set_message_handler(Arc::new(move |bytes: Bytes| {
let a_tx = a_tx.clone();
Box::pin(async move {
if let Some(tx) = a_tx.lock().await.take() {
let _ = tx.send(bytes.to_vec());
}
})
}))
.await;
node_b
.set_message_handler(Arc::new(move |bytes: Bytes| {
let b_tx = b_tx.clone();
Box::pin(async move {
if let Some(tx) = b_tx.lock().await.take() {
let _ = tx.send(bytes.to_vec());
}
})
}))
.await;
node_a.send(b"hello from a").await.unwrap();
node_b.send(b"hello from b").await.unwrap();
let received_by_b = tokio::time::timeout(Duration::from_secs(2), b_rx)
.await
.unwrap()
.unwrap();
let received_by_a = tokio::time::timeout(Duration::from_secs(2), a_rx)
.await
.unwrap()
.unwrap();
assert_eq!(received_by_b, b"hello from a");
assert_eq!(received_by_a, b"hello from b");
relay_task.abort();
});
}
#[test]
fn connected_without_subscription_returns_fallback_error() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let server_config = ServerConfig::builder()
.with_bind_default(0)
.with_identity(Identity::self_signed(["localhost"]).unwrap())
.keep_alive_interval(Some(Duration::from_secs(3)))
.build();
let server = Endpoint::server(server_config).unwrap();
let addr = server.local_addr().unwrap();
let server_task = tokio::spawn(async move {
let incoming = server.accept().await;
let request = incoming.await.unwrap();
let connection = request.accept().await.unwrap();
let (mut send, mut recv) = connection.accept_bi().await.unwrap();
let mut control_buf = Vec::new();
let mut chunk = [0u8; 256];
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
let setup = try_decode_control_frame(&mut control_buf).unwrap().unwrap();
assert_eq!(setup.0, MSG_CLIENT_SETUP);
send.write_all(&encode_server_setup_for_tests(
MOQ_DRAFT_07_VERSION,
SETUP_ROLE_BOTH,
))
.await
.unwrap();
loop {
if let Some(frame) = try_decode_control_frame(&mut control_buf).unwrap() {
assert_eq!(frame.0, MSG_ANNOUNCE);
break;
}
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
}
send.write_all(&encode_announce_ok_for_tests("local"))
.await
.unwrap();
let mut observed_subscribes = 0usize;
while observed_subscribes < 2 {
if let Some(frame) = try_decode_control_frame(&mut control_buf).unwrap() {
assert_eq!(frame.0, MSG_SUBSCRIBE);
observed_subscribes += 1;
continue;
}
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
}
});
let session = NativeMoQSession::new(
"local",
"remote",
&make_config(&format!("https://localhost:{}/moq", addr.port())),
)
.await
.unwrap();
session.start().await.unwrap();
assert_eq!(session.state(), NativeMoQState::Connected);
let err = session.send(b"hello").await.unwrap_err();
assert!(
err.to_string().contains("not subscribed"),
"expected not-subscribed fallback error, got: {err}"
);
server_task.await.unwrap();
});
}
#[test]
fn inbound_object_stream_buffers_until_handler_is_registered() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let server_config = ServerConfig::builder()
.with_bind_default(0)
.with_identity(Identity::self_signed(["localhost"]).unwrap())
.keep_alive_interval(Some(Duration::from_secs(3)))
.build();
let server = Endpoint::server(server_config).unwrap();
let addr = server.local_addr().unwrap();
let server_task = tokio::spawn(async move {
let incoming = server.accept().await;
let request = incoming.await.unwrap();
let connection = request.accept().await.unwrap();
let (mut send, mut recv) = connection.accept_bi().await.unwrap();
let mut control_buf = Vec::new();
let mut chunk = [0u8; 256];
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
let setup = try_decode_control_frame(&mut control_buf).unwrap().unwrap();
assert_eq!(setup.0, MSG_CLIENT_SETUP);
send.write_all(&encode_server_setup_for_tests(
MOQ_DRAFT_07_VERSION,
SETUP_ROLE_BOTH,
))
.await
.unwrap();
loop {
if let Some(frame) = try_decode_control_frame(&mut control_buf).unwrap() {
assert_eq!(frame.0, MSG_ANNOUNCE);
break;
}
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
}
send.write_all(&encode_announce_ok_for_tests("local"))
.await
.unwrap();
let mut observed_subscribes = 0usize;
while observed_subscribes < 2 {
if let Some(frame) = try_decode_control_frame(&mut control_buf).unwrap() {
assert_eq!(frame.0, MSG_SUBSCRIBE);
observed_subscribes += 1;
continue;
}
let read = recv.read(&mut chunk).await.unwrap().unwrap();
control_buf.extend_from_slice(&chunk[..read]);
}
let mut object_stream = connection.open_uni().await.unwrap().await.unwrap();
object_stream
.write_all(&encode_subgroup_object(1, 0, 1, 0, b"inbound"))
.await
.unwrap();
object_stream.finish().await.unwrap();
});
let session = NativeMoQSession::new(
"local",
"remote",
&make_config(&format!("https://localhost:{}/moq", addr.port())),
)
.await
.unwrap();
session.start().await.unwrap();
assert_eq!(session.state(), NativeMoQState::Connected);
tokio::time::sleep(Duration::from_millis(75)).await;
let (tx, rx) = tokio::sync::oneshot::channel::<Vec<u8>>();
let tx = Arc::new(Mutex::new(Some(tx)));
session
.set_message_handler(Arc::new(move |bytes: Bytes| {
let tx = tx.clone();
Box::pin(async move {
if let Some(tx) = tx.lock().await.take() {
let _ = tx.send(bytes.to_vec());
}
})
}))
.await;
let received = tokio::time::timeout(Duration::from_secs(2), rx)
.await
.unwrap()
.unwrap();
assert_eq!(received, b"inbound");
server_task.await.unwrap();
});
}
#[test]
fn invalid_relay_url_reaches_failed() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let session =
NativeMoQSession::new("local", "remote", &make_config("https://relay.invalid/moq"))
.await
.unwrap();
let result = session.start().await;
assert!(
result.is_err(),
"expected start() to fail for .invalid relay"
);
assert_eq!(session.state(), NativeMoQState::Failed);
});
}
#[test]
fn empty_relay_url_is_rejected_at_construction() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let result = NativeMoQSession::new("local", "remote", &make_config("")).await;
assert!(
result.is_err(),
"expected new() to fail for empty relay_url"
);
});
}
#[test]
fn send_without_connected_session_returns_not_started() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let session = NativeMoQSession::new(
"local",
"remote",
&make_config("https://relay.example.com/moq"),
)
.await
.unwrap();
let err = session.send(b"hello").await.unwrap_err();
assert!(
err.to_string().contains("not started"),
"expected not-started error, got: {err}"
);
});
}
#[test]
fn close_reaches_closed_state() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async {
let session = NativeMoQSession::new(
"local",
"remote",
&make_config("https://relay.example.com/moq"),
)
.await
.unwrap();
session.close();
assert_eq!(session.state(), NativeMoQState::Closed);
});
}
}